diff --git a/sdks/python/apache_beam/dataframe/frames.py b/sdks/python/apache_beam/dataframe/frames.py index 310791d2b58f..17bd04ba1338 100644 --- a/sdks/python/apache_beam/dataframe/frames.py +++ b/sdks/python/apache_beam/dataframe/frames.py @@ -1091,27 +1091,88 @@ def xs(self, key, axis, level, **kwargs): reindexed = self.reorder_levels( level + [i for i in range(self.index.nlevels) if i not in level]) - def xs_partitioned(frame, key): - if not len(key): - # key is not in this partition, return empty dataframe - result = frame.iloc[:0] - if key_size < frame.index.nlevels: + if key_size < reindexed.index.nlevels: + + def xs_partitioned(frame, key): + if not len(key): + # key is not in this partition, return empty dataframe/series + result = frame.iloc[:0] return result.droplevel(list(range(key_size))) - else: - return result + return frame.xs(key.item(), **kwargs) - # key should be in this partition, call xs. Will raise KeyError if not - # present. - return frame.xs(key.item()) + return frame_base.DeferredFrame.wrap( + expressions.ComputedExpression( + 'xs', + xs_partitioned, [reindexed._expr, key_expr], + requires_partition_by=partitionings.Index(list(range(key_size))), + preserves_partition_by=partitionings.Singleton())) + else: + # When all index levels are matched (key_size >= nlevels), pandas .xs() + # return type is data-dependent: + # - Single match: reduces dimensionality (DataFrame -> Series, Series -> scalar) + # - Duplicate matches: preserves container type (DataFrame -> DataFrame, Series -> Series) + # Because proxy schemas are 0-row templates evaluated at graph construction time + # without knowledge of dataset contents or key frequencies, the proxy always assumes + # a single match (dimensionality-reduced type). At runtime, the Singleton unwrap stage + # correctly produces whichever type pandas returns. Tests with multi-matching keys + # therefore specify check_proxy=False. + def xs_partitioned_wrapped(frame, key): + if not len(key): + return pd.Series([], dtype=object) + k = key.item() + try: + res = frame.xs(k, **kwargs) + return pd.Series([res], dtype=object) + except KeyError: + return pd.Series([], dtype=object) + + intermediate = expressions.ComputedExpression( + 'xs_partitioned_wrapped', + xs_partitioned_wrapped, [reindexed._expr, key_expr], + proxy=pd.Series([], dtype=object), + requires_partition_by=partitionings.Index(list(range(key_size))), + preserves_partition_by=partitionings.Singleton()) - return frame_base.DeferredFrame.wrap( - expressions.ComputedExpression( - 'xs', - xs_partitioned, - [reindexed._expr, key_expr], - requires_partition_by=partitionings.Index(list(range(key_size))), - # Drops index levels, so partitioning is not preserved - preserves_partition_by=partitionings.Singleton())) + proxy_frame = reindexed._expr.proxy() + k_val = key_series.iloc[0] + if isinstance(proxy_frame, pd.DataFrame): + if not proxy_frame.index.is_unique: + proxy_frame = proxy_frame.copy() + proxy_frame.index = proxy_frame.index.drop_duplicates() + dummy_index = ( + pd.MultiIndex.from_tuples([k_val], names=proxy_frame.index.names) if + isinstance(k_val, tuple) else pd.Index([k_val], + name=proxy_frame.index.name)) + dummy_obj = proxy_frame.reindex(dummy_index) + xs_proxy = dummy_obj.xs(k_val, **kwargs) + if isinstance(xs_proxy, (pd.DataFrame, pd.Series)): + xs_proxy = xs_proxy.iloc[:0] + else: + try: + xs_proxy = proxy_frame.dtype.type() + except TypeError: + dummy_index = ( + pd.MultiIndex.from_tuples([k_val], names=proxy_frame.index.names) + if isinstance(k_val, tuple) else pd.Index( + [k_val], name=proxy_frame.index.name)) + if not proxy_frame.index.is_unique: + proxy_frame = proxy_frame.copy() + proxy_frame.index = proxy_frame.index.drop_duplicates() + xs_proxy = proxy_frame.reindex(dummy_index).iloc[0] + + def unwrap_xs(ser): + if ser.empty: + raise KeyError(k_val) + return ser.iloc[0] + + with expressions.allow_non_parallel_operations(True): + return frame_base.DeferredFrame.wrap( + expressions.ComputedExpression( + 'xs', + unwrap_xs, [intermediate], + proxy=xs_proxy, + requires_partition_by=partitionings.Singleton(), + preserves_partition_by=partitionings.Singleton())) @property def dtype(self): diff --git a/sdks/python/apache_beam/dataframe/frames_test.py b/sdks/python/apache_beam/dataframe/frames_test.py index 7a03af6220b8..a2cb2e498005 100644 --- a/sdks/python/apache_beam/dataframe/frames_test.py +++ b/sdks/python/apache_beam/dataframe/frames_test.py @@ -331,6 +331,19 @@ def test_series_xs(self): lambda df: df.num_legs.xs(('bird', 'walks'), level=[0, 'locomotion']), df) + # Test cases reported in BEAM-28559 + df_single_index = df.reset_index().set_index('class') + self._run_test( + lambda df: df.num_legs.xs('mammal'), df_single_index, check_proxy=False) + self._run_test(lambda df: df.num_legs.xs('bird'), df_single_index) + + # Categorical Series single match + s_cat = pd.Series( + pd.Categorical(['a', 'b', 'c']), + index=['r1', 'r2', 'r3'], + name='cat_col') + self._run_test(lambda s: s.xs('r1'), s_cat, check_proxy=False) + def test_dataframe_xs(self): # Test cases reported in BEAM-13421 df = pd.DataFrame( @@ -342,10 +355,53 @@ def test_dataframe_xs(self): ]), columns=['provider', 'time', 'value']) - self._run_test(lambda df: df.xs('state'), df.set_index(['provider'])) + self._run_test( + lambda df: df.xs('state'), + df.set_index(['provider']), + check_proxy=False) self._run_test( lambda df: df.xs('state'), df.set_index(['provider', 'time'])) + # Test cases reported in BEAM-28559 + self._run_test(lambda df: df.xs('county'), df.set_index(['provider'])) + self._run_test( + lambda df: df.xs(('state', 'day1')), + df.set_index(['provider', 'time']), + check_proxy=False) + + df_unique = pd.DataFrame( + np.array([ + ['state', 'day1', 12], + ['state', 'day2', 14], + ['county', 'day1', 9], + ]), + columns=['provider', 'time', 'value']) + self._run_test( + lambda df: df.xs(('state', 'day2')), + df_unique.set_index(['provider', 'time'])) + + # Categorical and extension dtype tests + df_cat = pd.DataFrame({ + 'cat': pd.Categorical(['a', 'b', 'c']), 'val': [1, 2, 3] + }, + index=['r1', 'r2', 'r3']) + self._run_test(lambda df: df.xs('r1'), df_cat) + + df_dt_tz = pd.DataFrame({ + 'dt': pd.Series([ + pd.Timestamp('2023-01-01', tz='UTC'), + pd.Timestamp('2023-01-02', tz='UTC') + ], + dtype='datetime64[ns, UTC]'), + 'val': [1, 2] + }, + index=['r1', 'r2']) + self._run_test(lambda df: df.xs('r1'), df_dt_tz) + + df_null_int = pd.DataFrame({'num': pd.Series([1, 2, None], dtype='Int64')}, + index=['r1', 'r2', 'r3']) + self._run_test(lambda df: df.xs('r1'), df_null_int) + def test_set_column(self): def new_column(df): df['NewCol'] = df['Speed']