diff --git a/python/cudf_polars/cudf_polars/dsl/expressions/rolling.py b/python/cudf_polars/cudf_polars/dsl/expressions/rolling.py index 15a2f26e5255..0cb0d8f06e59 100644 --- a/python/cudf_polars/cudf_polars/dsl/expressions/rolling.py +++ b/python/cudf_polars/cudf_polars/dsl/expressions/rolling.py @@ -62,6 +62,11 @@ class CumSumOp(UnaryOp): pass +@dataclass(frozen=True) +class ShiftOp(UnaryOp): + pass + + def to_request( value: expr.Expr, orderby: Column, df: DataFrame ) -> plc.rolling.RollingRequest: @@ -359,7 +364,13 @@ def __init__( or ( isinstance(named_expr.value, expr.UnaryFunction) and named_expr.value.name - in {"rank", "fill_null_with_strategy", "cum_sum"} + in { + "rank", + "fill_null_with_strategy", + "cum_sum", + "shift", + "shift_and_fill", + } ) ) ] @@ -569,6 +580,64 @@ def _( # type: ignore[no-untyped-def] result_tables.append(filled) return out_names, out_dtypes, result_tables + @_apply_unary_op.register + def _( + self, + op: ShiftOp, + df: DataFrame, + _: plc.groupby.GroupBy, + ) -> tuple[list[str], list[DataType], list[plc.Table]]: + plc_cols: list[plc.Column] = [] + offsets: list[int] = [] + fill_scalars: list[plc.Scalar] = [] + out_names: list[str] = [] + out_dtypes: list[DataType] = [] + + for ne in op.named_exprs: + shift_expr = ne.value + assert isinstance(shift_expr, expr.UnaryFunction) + data_expr, offset_expr = shift_expr.children[:2] + assert isinstance(offset_expr, expr.Literal) + offset = offset_expr.value + assert isinstance(offset, int) + + plc_col = data_expr.evaluate(df, context=ExecutionContext.FRAME).obj + plc_cols.append(plc_col) + offsets.append(offset) + out_names.append(ne.name) + out_dtypes.append(shift_expr.dtype) + if shift_expr.name == "shift": + fill_scalars.append( + plc.Scalar.from_py(None, plc_col.type(), stream=df.stream) + ) + else: + assert shift_expr.name == "shift_and_fill" + fill_expr = shift_expr.children[2] + assert isinstance(fill_expr, expr.Literal) + fill_scalars.append( + plc.Scalar.from_py( + fill_expr.value, plc_col.type(), stream=df.stream + ) + ) + + assert op.order_index is not None + val_cols = plc.copying.gather( + plc.Table(plc_cols), + op.order_index, + plc.copying.OutOfBoundsPolicy.NULLIFY, + stream=df.stream, + ).columns() + + assert isinstance(op.local_grouper, plc.groupby.GroupBy) + shifted_tbl = op.local_grouper.shift( + plc.Table(val_cols), offsets, fill_scalars, stream=df.stream + )[1] + return ( + out_names, + out_dtypes, + [plc.Table([column]) for column in shifted_tbl.columns()], + ) + def _reorder_to_input( self, row_id: plc.Column, @@ -627,6 +696,7 @@ def _split_named_expr( "rank": [], "fill_null_with_strategy": [], "cum_sum": [], + "shift": [], } for ne in self.named_aggs: @@ -640,6 +710,8 @@ def _split_named_expr( unary_window_ops["cum_sum"].append(ne) elif isinstance(v, expr.UnaryFunction) and v.name in unary_window_ops: unary_window_ops[v.name].append(ne) + elif isinstance(v, expr.UnaryFunction) and v.name == "shift_and_fill": + unary_window_ops["shift"].append(ne) else: reductions.append(ne) return reductions, unary_window_ops @@ -1119,6 +1191,48 @@ def do_evaluate( # noqa: D102 ) ) + if shift_named := unary_window_ops["shift"]: + order_index, shift_by_cols_for_scan, local = ( + self._grouped_window_scan_setup( + by_cols, + row_id=row_id, + order_by_col=order_by_col + if self._order_by_expr is not None + else None, + ob_desc=self.options[2] + if self._order_by_expr is not None + else False, + ob_nulls_last=self.options[3] + if self._order_by_expr is not None + else False, + grouper=grouper, + stream=df.stream, + require_sorted_groups=True, + ) + ) + names, dtypes, tables = self._apply_unary_op( + ShiftOp( + named_exprs=shift_named, + order_index=order_index, + by_cols_for_scan=shift_by_cols_for_scan, + local_grouper=local, + ), + df, + grouper, + ) + broadcasted_cols.extend( + self._reorder_to_input( + row_id, + by_cols, + df.num_rows, + tables, + names, + dtypes, + order_index=order_index, + stream=df.stream, + ) + ) + # Create a temporary DataFrame with the broadcasted columns named by their # placeholder names from agg decomposition, then evaluate the post-expression. df = DataFrame(broadcasted_cols, stream=df.stream) diff --git a/python/cudf_polars/cudf_polars/dsl/translate.py b/python/cudf_polars/cudf_polars/dsl/translate.py index 8288a643226e..8316879b0c12 100644 --- a/python/cudf_polars/cudf_polars/dsl/translate.py +++ b/python/cudf_polars/cudf_polars/dsl/translate.py @@ -129,7 +129,8 @@ def _unsupported_fill_over_window(value: expr.Expr) -> bool: windowed = [ node for node in traversal([value]) - if isinstance(node, expr.UnaryFunction) and node.name in {"rank", "cum_sum"} + if isinstance(node, expr.UnaryFunction) + and node.name in {"rank", "cum_sum", "shift", "shift_and_fill"} ] if not windowed: return False @@ -1229,7 +1230,14 @@ def _( if isinstance(v, expr.Agg) or ( isinstance(v, expr.UnaryFunction) - and v.name in {"rank", "fill_null_with_strategy", "cum_sum"} + and v.name + in { + "rank", + "fill_null_with_strategy", + "cum_sum", + "shift", + "shift_and_fill", + } ) ] children = (*by_exprs, *((order_by_expr,) if has_order_by else ()), *child_deps) diff --git a/python/cudf_polars/cudf_polars/dsl/utils/aggregations.py b/python/cudf_polars/cudf_polars/dsl/utils/aggregations.py index b595a779f87b..ded67a6cfabd 100644 --- a/python/cudf_polars/cudf_polars/dsl/utils/aggregations.py +++ b/python/cudf_polars/cudf_polars/dsl/utils/aggregations.py @@ -98,11 +98,24 @@ def decompose_single_agg( "rank", "fill_null_with_strategy", "cum_sum", + "shift", + "shift_and_fill", }: if context != ExecutionContext.WINDOW: raise NotImplementedError( f"{agg.name} is not supported in groupby or rolling context" ) + if agg.name in {"shift", "shift_and_fill"}: + if not isinstance(agg.children[1], expr.Literal): + raise NotImplementedError( + "shift over a window only supports a literal offset" + ) + if agg.name == "shift_and_fill" and not isinstance( + agg.children[2], expr.Literal + ): + raise NotImplementedError( + "shift over a window only supports a literal fill_value" + ) if agg.name == "fill_null_with_strategy" and ( strategy := agg.options[0] ) not in {"forward", "backward"}: diff --git a/python/cudf_polars/cudf_polars/streaming/select.py b/python/cudf_polars/cudf_polars/streaming/select.py index 223f1745587a..53196c562e18 100644 --- a/python/cudf_polars/cudf_polars/streaming/select.py +++ b/python/cudf_polars/cudf_polars/streaming/select.py @@ -25,7 +25,7 @@ from cudf_polars.streaming.over import _fuse_over_nodes from cudf_polars.streaming.repartition import Repartition from cudf_polars.streaming.utils import ( - _contains_cum_sum_without_order_by, + _contains_input_order_window_without_order_by, _contains_unsupported_fill_strategy, _dynamic_planning_on, _lower_ir_fallback, @@ -414,15 +414,16 @@ def _( ), ) - if rec.state["nranks"] > 1 and _contains_cum_sum_without_order_by( + if rec.state["nranks"] > 1 and _contains_input_order_window_without_order_by( [e.value for e in ir.exprs] ): return _lower_ir_fallback( ir.reconstruct([child]), rec, msg=( - "cum_sum() over a window without order_by is not supported across " - "multiple ranks; falling back to a single partition." + "input-order-sensitive window expressions without order_by are " + "not supported across multiple ranks; falling back to a single " + "partition." ), ) diff --git a/python/cudf_polars/cudf_polars/streaming/utils.py b/python/cudf_polars/cudf_polars/streaming/utils.py index d797b588ed2c..b1e50301832b 100644 --- a/python/cudf_polars/cudf_polars/streaming/utils.py +++ b/python/cudf_polars/cudf_polars/streaming/utils.py @@ -125,8 +125,11 @@ def _contains_unsupported_fill_strategy(exprs: Sequence[Expr]) -> bool: return False -def _contains_cum_sum_without_order_by(exprs: Sequence[Expr]) -> bool: - # Returns True for cum_sum(...) or fill_null_with_strategy(cum_sum(...)). +_INPUT_ORDER_WINDOW_OPS = frozenset({"cum_sum", "shift", "shift_and_fill"}) + + +def _contains_input_order_window_without_order_by(exprs: Sequence[Expr]) -> bool: + """Return True for implicit input-order-sensitive window expressions.""" for e in traversal(exprs): if not (isinstance(e, GroupedWindow) and not e.options[1]): continue @@ -138,6 +141,6 @@ def _contains_cum_sum_without_order_by(exprs: Sequence[Expr]) -> bool: and isinstance(v.children[0], UnaryFunction) ): v = v.children[0] - if isinstance(v, UnaryFunction) and v.name == "cum_sum": + if isinstance(v, UnaryFunction) and v.name in _INPUT_ORDER_WINDOW_OPS: return True return False diff --git a/python/cudf_polars/tests/expressions/test_rolling.py b/python/cudf_polars/tests/expressions/test_rolling.py index 570497147d99..c579e22abde6 100644 --- a/python/cudf_polars/tests/expressions/test_rolling.py +++ b/python/cudf_polars/tests/expressions/test_rolling.py @@ -441,6 +441,66 @@ def test_cum_sum_over( assert_gpu_result_equal(q, engine=engine) +@pytest.mark.parametrize("n", [1, -1, 2]) +@pytest.mark.parametrize( + "expr,group_key", + [ + (pl.col("x"), "g"), + (pl.when((pl.col("x") % 2) == 0).then(None).otherwise(pl.col("x")), "g"), + (pl.col("x"), "g_null"), + ], +) +@pytest.mark.parametrize("order_by", ["x2", ["g2", pl.col("x2") * 2]]) +def test_shift_over( + engine: pl.GPUEngine, + df: pl.LazyFrame, + n: int, + expr: pl.Expr, + group_key: str, + order_by: str | list[str | pl.Expr], +) -> None: + q = df.select(expr.shift(n).over(group_key, order_by=order_by)) + assert_gpu_result_equal(q, engine=engine) + + +@pytest.mark.parametrize("n,fill_value", [(1, 0), (-1, 99)]) +def test_shift_over_fill_value( + engine: pl.GPUEngine, + df: pl.LazyFrame, + n: int, + fill_value: int, +) -> None: + q = df.select(pl.col("x").shift(n, fill_value=fill_value).over("g", order_by="x2")) + assert_gpu_result_equal(q, engine=engine) + + +@pytest.mark.parametrize( + "expr", + [ + pl.col("x").shift(pl.col("x2").min()).over("g"), + pl.col("x").shift(1, fill_value=pl.col("x2").min()).over("g"), + ], + ids=["nonliteral_offset", "nonliteral_fill_value"], +) +def test_shift_over_nonliteral_args_raises( + engine: pl.GPUEngine, + df: pl.LazyFrame, + expr: pl.Expr, +) -> None: + q = df.select(expr) + assert_ir_translation_raises(q, engine, NotImplementedError) + + +@pytest.mark.parametrize("n", [1, -1]) +def test_shift_over_without_order_by( + engine_raise_on_fail: pl.GPUEngine, + df: pl.LazyFrame, + n: int, +) -> None: + q = df.select(pl.col("x").shift(n).over("g")) + assert_gpu_result_equal(q, engine=engine_raise_on_fail) + + @pytest.mark.parametrize( "expr", [ diff --git a/python/cudf_polars/tests/streaming/test_rolling.py b/python/cudf_polars/tests/streaming/test_rolling.py index 8ebf6c751b02..fd6fd34c1e68 100644 --- a/python/cudf_polars/tests/streaming/test_rolling.py +++ b/python/cudf_polars/tests/streaming/test_rolling.py @@ -63,6 +63,8 @@ def test_rolling_datetime(engine): pl.col("x").rank(method="dense", descending=True).over("g"), pl.col("x").rank(method="min").over("g", "g2"), pl.col("x").cum_sum().over("g", order_by="s"), + pl.col("x").shift(1).over("g", order_by="s"), + pl.col("x").shift(-1, fill_value=0).over("g", order_by="s"), pl.when((pl.col("x") % 2) == 0) .then(None) .otherwise(pl.col("x")) @@ -78,6 +80,8 @@ def test_rolling_datetime(engine): "rank_dense", "rank_min_multi_key", "cum_sum_order_by", + "shift_order_by", + "shift_fill_order_by", "fill_null_forward", ], ) @@ -107,6 +111,20 @@ def test_over_cum_sum_fill_null_per_partition(engine, strategy): assert_gpu_result_equal(df.select(expr), engine=engine, check_row_order=True) +def test_over_shift_without_order_by_single_rank(spmd_engine_factory) -> None: + engine = spmd_engine_factory( + StreamingOptions(max_rows_per_partition=2, fallback_mode="raise"), + ) + df = pl.LazyFrame( + { + "g": [1, 1, 2, 2, 2, 1], + "x": [1, 2, 3, 4, 5, 6], + } + ) + q = df.select(pl.col("x").shift(1).over("g")) + assert_gpu_result_equal(q, engine=engine, check_row_order=True) + + @pytest.mark.parametrize( "expr", [ @@ -127,6 +145,7 @@ def test_over_cum_sum_fill_null_per_partition(engine, strategy): .fill_null(strategy="forward") .fill_null(strategy="forward") .over("g", order_by="x"), + pl.col("x").shift(1).fill_null(strategy="forward").over("g"), ], ids=[ "rank_fill", @@ -134,10 +153,11 @@ def test_over_cum_sum_fill_null_per_partition(engine, strategy): "rank_abs_fill", "cum_sum_abs_fill", "cum_sum_fill_fill", + "shift_fill", ], ) def test_over_fill_null_over_window_fails_translation(engine, expr): - df = pl.LazyFrame({"g": [1, 1, 2, 2, 2, 1], "x": [1.0, 2, 3, 4, 5, 6]}) + df = pl.LazyFrame({"g": [1, 1, 2, 2, 2, 1], "x": [1.0, None, 3, 4, None, 6]}) assert_ir_translation_raises(df.select(expr), engine, NotImplementedError) @@ -231,8 +251,15 @@ def test_over_mixed_keys(streaming_engine_factory) -> None: pl.len().over("g"), pl.col("x").rank(method="dense").over("g"), pl.col("x").cum_sum().over("g", order_by="s"), + pl.col("x").shift(1).over("g", order_by="s"), + ], + ids=[ + "scalar_sum", + "scalar_len", + "nonscalar_rank", + "nonscalar_cum_sum", + "nonscalar_shift", ], - ids=["scalar_sum", "scalar_len", "nonscalar_rank", "nonscalar_cum_sum"], ) @pytest.mark.parametrize("max_rows_per_partition", [1, 2]) def test_over_many_partitions( @@ -283,8 +310,9 @@ def test_over_empty_input(streaming_engine_factory, expr) -> None: [ pl.col("x").sum().over("g"), pl.col("x").rank(method="dense").over("g"), + pl.col("x").shift(1).over("g", order_by="x"), ], - ids=["scalar_sum", "nonscalar_rank"], + ids=["scalar_sum", "nonscalar_rank", "nonscalar_shift"], ) def test_over_already_partitioned(streaming_engine_factory, expr) -> None: # broadcast_limit=0 disables broadcast joins entirely. Therefore, we should diff --git a/python/cudf_polars/tests/streaming/test_spmd.py b/python/cudf_polars/tests/streaming/test_spmd.py index ad67a6252ab0..1d50107e663b 100644 --- a/python/cudf_polars/tests/streaming/test_spmd.py +++ b/python/cudf_polars/tests/streaming/test_spmd.py @@ -494,12 +494,13 @@ def test_quent_context_default(spmd_engine: SPMDEngine) -> None: @pytest.mark.parametrize( - "expr,is_scalar", + "expr,expected", [ - (pl.col("x").sum().over("g").alias("result"), True), - (pl.col("x").rank(method="dense").over("g").alias("result"), False), + (pl.col("x").sum().over("g").alias("result"), "sum"), + (pl.col("x").rank(method="dense").over("g").alias("result"), "rank"), + (pl.col("x").shift(1).over("g", order_by="x").alias("result"), "shift"), ], - ids=["scalar_sum", "nonscalar_rank"], + ids=["scalar_sum", "nonscalar_rank", "nonscalar_shift"], ) @pytest.mark.parametrize( "cross_rank", @@ -507,10 +508,9 @@ def test_quent_context_default(spmd_engine: SPMDEngine) -> None: ids=["same_rank", "cross_rank"], ) def test_over_multirank( - request: pytest.FixtureRequest, comm: Communicator, expr: pl.Expr, - is_scalar: bool, # noqa: FBT001 + expected: str, cross_rank: bool, # noqa: FBT001 ) -> None: """over() correctness in multi-rank SPMD mode, same-rank and cross-rank cases. @@ -530,11 +530,7 @@ def test_over_multirank( rank = engine.rank nranks = engine.nranks if nranks != 2: - request.applymarker( - pytest.mark.skip( - reason="key assignments are probed for exactly 2 ranks" - ) - ) + pytest.skip("key assignments are probed for exactly 2 ranks") keys = _CROSS_RANK_KEYS if cross_rank else _SAME_RANK_KEYS g = keys[rank] xs = [rank * 3 + 1, rank * 3 + 2, rank * 3 + 3] @@ -559,14 +555,42 @@ def test_over_multirank( assert grp.shape == (3, 3), f"rank {r} group has wrong row count" expected_xs = [r * 3 + 1, r * 3 + 2, r * 3 + 3] assert grp["x"].to_list() == expected_xs - if is_scalar: + if expected == "sum": assert grp["result"].to_list() == [sum(expected_xs)] * 3 - else: + elif expected == "rank": assert grp["result"].to_list() == [1, 2, 3] + else: + assert grp["result"].to_list() == [None, *expected_xs[:-1]] + + +def test_over_shift_without_order_by_multirank_raises(comm: Communicator) -> None: + with SPMDEngine( + comm=comm, + executor_options={ + "max_rows_per_partition": 2, + "dynamic_planning": {}, + "fallback_mode": "raise", + }, + ) as engine: + if engine.nranks < 2: + pytest.skip("requires multiple ranks") + + rank = engine.rank + lf = pl.LazyFrame( + { + "g": [0, 0, 0], + "x": [rank * 3 + 1, rank * 3 + 2, rank * 3 + 3], + } + ) + q = lf.select(pl.col("x").shift(1).over("g")) + with pytest.raises( + NotImplementedError, + match=r"input-order-sensitive window expressions without order_by", + ): + q.collect(engine=engine) def test_over_nonscalar_duplicated_input( - request: pytest.FixtureRequest, comm: Communicator, ) -> None: """Non-scalar over() on duplicated=True input produces correct row count and values. @@ -586,11 +610,7 @@ def test_over_nonscalar_duplicated_input( rank = engine.rank nranks = engine.nranks if nranks != 2: - request.applymarker( - pytest.mark.skip( - reason="key assignments are probed for exactly 2 ranks" - ) - ) + pytest.skip("key assignments are probed for exactly 2 ranks") coarse_g = _SAME_RANK_KEYS[rank] fine_gs = [rank * 3 + 1, rank * 3 + 2, rank * 3 + 3]