diff --git a/python/cudf_polars/cudf_polars/streaming/partitioning_requests.py b/python/cudf_polars/cudf_polars/streaming/partitioning_requests.py new file mode 100644 index 00000000000..aaf6e79ff35 --- /dev/null +++ b/python/cudf_polars/cudf_polars/streaming/partitioning_requests.py @@ -0,0 +1,276 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +""" +Downstream partitioning requests for streaming actor-graph construction. + +A partitioning request is attached to an IR node when a downstream consumer +may benefit from that node producing a specific partitioning. These requests +use "partitioning" in the same broad sense as ``ChannelMetadata.Partitioning``: +rows may be strictly partitioned by equality keys, ordered by key values, or +both. + +Requests are planning-time information. They do not describe or guarantee the +actual partitioning of the node's output. Runtime partitioning metadata is +tracked separately in ``ChannelMetadata``. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING, TypeAlias + +import pylibcudf as plc + +from cudf_polars.dsl import expr +from cudf_polars.dsl.ir import ( + Filter, + GroupBy, + Join, + MapFunction, + Projection, + Select, + Slice, + Sort, +) +from cudf_polars.dsl.traversal import post_traversal +from cudf_polars.dsl.utils.column_domain import column_domain_bindings + +if TYPE_CHECKING: + from collections.abc import Mapping + + from cudf_polars.dsl.ir import IR + + +@dataclass(frozen=True) +class NamedOrderKey: + """Named sort key with logical Polars ordering options.""" + + name: str + descending: bool + nulls_last: bool + + +@dataclass(frozen=True) +class StrictPartitioningRequest: + """Request for upstream output to strictly partition equal keys.""" + + keys: tuple[str, ...] + + +@dataclass(frozen=True) +class OrderPartitioningRequest: + """Request for upstream output to be ordered by a key sequence.""" + + keys: tuple[NamedOrderKey, ...] + strict_key_count: int | None = None + + +PartitioningRequest: TypeAlias = StrictPartitioningRequest | OrderPartitioningRequest + + +def collect_partitioning_requests( + ir: IR, +) -> dict[IR, tuple[PartitioningRequest, ...]]: + """ + Collect downstream partitioning requests for each IR node. + + The returned mapping answers "which partitionings could make downstream + consumers cheaper if this node produced them?" A request is therefore + aspirational: it is not evidence that the data is currently sorted, hash + partitioned, or otherwise partitioned that way. + """ + requests: dict[IR, tuple[PartitioningRequest, ...]] = {} + # Reverse post-order ensures every downstream consumer is processed before + # any shared upstream producer whose requests must be propagated further. + for node in reversed(list(post_traversal([ir]))): + child_requests = _direct_child_requests(node) + child_requests.extend(_propagated_child_requests(node, requests)) + for child, request in child_requests: + requests[child] = _merge_candidate_request(requests.get(child, ()), request) + return requests + + +def _direct_child_requests(ir: IR) -> list[tuple[IR, PartitioningRequest]]: + """Create child requests implied by partitioning-aware operators.""" + if isinstance(ir, Sort): + names = _column_names(ir.by) + if names is not None: + return [ + ( + ir.children[0], + _order_request( + names, + tuple( + order == plc.types.Order.DESCENDING for order in ir.order + ), + tuple( + (order == plc.types.Order.ASCENDING) + == (null_order == plc.types.NullOrder.AFTER) + for order, null_order in zip( + ir.order, ir.null_order, strict=True + ) + ), + ), + ) + ] + + if isinstance(ir, MapFunction) and ir.name == "hint_sorted": + return [(ir.children[0], _order_request(*ir.options))] + + if isinstance(ir, Join) and ir.options[0] != "Cross": + left_keys = _column_names(ir.left_on) + right_keys = _column_names(ir.right_on) + if left_keys is not None and right_keys is not None: + return [ + (ir.children[0], StrictPartitioningRequest(left_keys)), + (ir.children[1], StrictPartitioningRequest(right_keys)), + ] + + if isinstance(ir, GroupBy) and not ir.maintain_order: + keys = _column_names(ir.keys) + if keys is not None: + return [(ir.children[0], StrictPartitioningRequest(keys))] + + return [] + + +def _propagated_child_requests( + node: IR, requests: dict[IR, tuple[PartitioningRequest, ...]] +) -> list[tuple[IR, PartitioningRequest]]: + """Push compatible downstream requests through single-child operators.""" + child_requests: list[tuple[IR, PartitioningRequest]] = [] + node_requests = requests.get(node) + if ( + node_requests is not None + and len(node.children) == 1 + and isinstance(node, (Projection, Select, Filter, Slice, GroupBy)) + ): + remapping = { + output_name: binding.name + for output_name, binding in column_domain_bindings(node).items() + if binding.child_index == 0 + } + child_requests.extend( + (node.children[0], remapped) + for node_request in node_requests + if (remapped := _remap_request(node_request, remapping)) is not None + ) + return child_requests + + +def _merge_candidate_request( + existing_requests: tuple[PartitioningRequest, ...], + request: PartitioningRequest, +) -> tuple[PartitioningRequest, ...]: + """Merge compatible requests while preserving incompatible candidates.""" + candidates: list[PartitioningRequest] = [] + new_request = request + insertion_index: int | None = None + for existing_request in existing_requests: + if (merged := _merge_requests(existing_request, new_request)) is None: + candidates.append(existing_request) + else: + new_request = merged + if insertion_index is None: + insertion_index = len(candidates) + if insertion_index is None: + candidates.append(new_request) + else: + candidates.insert(insertion_index, new_request) + return tuple(candidates) + + +def _order_request( + names: tuple[str, ...], + descending: tuple[bool, ...], + nulls_last: tuple[bool, ...], +) -> OrderPartitioningRequest: + return OrderPartitioningRequest( + tuple( + NamedOrderKey(name, desc, null_last) + for name, desc, null_last in zip(names, descending, nulls_last, strict=True) + ) + ) + + +def _column_names(named_exprs: tuple[expr.NamedExpr, ...]) -> tuple[str, ...] | None: + names = [] + for named_expr in named_exprs: + if not isinstance(named_expr.value, expr.Col): + return None + names.append(named_expr.value.name) + return tuple(names) + + +def _remap_request( + request: PartitioningRequest, remapping: Mapping[str, str] +) -> PartitioningRequest | None: + """Rewrite request column names through a child-to-parent name mapping.""" + if isinstance(request, StrictPartitioningRequest): + remapped_names = [] + for name in request.keys: + if (new_name := remapping.get(name)) is None: + return None + remapped_names.append(new_name) + return StrictPartitioningRequest(tuple(remapped_names)) + + remapped_keys = [] + for key in request.keys: + new_name = remapping.get(key.name) + if new_name is None: + return None + remapped_keys.append(NamedOrderKey(new_name, key.descending, key.nulls_last)) + return OrderPartitioningRequest(tuple(remapped_keys), request.strict_key_count) + + +def _merge_requests( + left: PartitioningRequest, right: PartitioningRequest +) -> PartitioningRequest | None: + """Merge compatible requests, or keep both candidates if incompatible.""" + if isinstance(left, StrictPartitioningRequest): + if isinstance(right, StrictPartitioningRequest): + if _is_prefix(left.keys, right.keys): + return left + if _is_prefix(right.keys, left.keys): + return right + return None + return _merge_order_with_strict(right, left) + + if isinstance(right, StrictPartitioningRequest): + return _merge_order_with_strict(left, right) + + if _is_prefix(left.keys, right.keys): + keys = right.keys + elif _is_prefix(right.keys, left.keys): + keys = left.keys + else: + return None + return OrderPartitioningRequest( + keys, _merge_strict_key_count(left.strict_key_count, right.strict_key_count) + ) + + +def _merge_order_with_strict( + order_request: OrderPartitioningRequest, + strict_request: StrictPartitioningRequest, +) -> OrderPartitioningRequest | None: + """Fold strict-key requirements into compatible ordering requests.""" + order_names = tuple(key.name for key in order_request.keys) + if _is_prefix(strict_request.keys, order_names) or _is_prefix( + order_names, strict_request.keys + ): + strict_key_count = min(len(strict_request.keys), len(order_names)) + return OrderPartitioningRequest( + order_request.keys, + _merge_strict_key_count(order_request.strict_key_count, strict_key_count), + ) + return None + + +def _merge_strict_key_count(*counts: int | None) -> int | None: + count = max((count for count in counts if count is not None), default=0) + return count or None + + +def _is_prefix(left: tuple[object, ...], right: tuple[object, ...]) -> bool: + return len(left) <= len(right) and left == right[: len(left)] diff --git a/python/cudf_polars/tests/streaming/test_partitioning_requests.py b/python/cudf_polars/tests/streaming/test_partitioning_requests.py new file mode 100644 index 00000000000..77d5256f408 --- /dev/null +++ b/python/cudf_polars/tests/streaming/test_partitioning_requests.py @@ -0,0 +1,523 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import polars as pl + +from cudf_polars.containers import DataType +from cudf_polars.dsl import expr +from cudf_polars.dsl.ir import ( + DataFrameScan, + GroupBy, + Join, + MapFunction, + Select, + Sort, + Union, +) +from cudf_polars.streaming.partitioning_requests import ( + NamedOrderKey, + OrderPartitioningRequest, + StrictPartitioningRequest, + collect_partitioning_requests, +) +from cudf_polars.utils.sorting import sort_order + +if TYPE_CHECKING: + from cudf_polars.dsl.ir import IR + +I64 = DataType(pl.Int64()) + + +def make_scan(*names: str) -> DataFrameScan: + frame = pl.DataFrame({name: [1] for name in names}) + return DataFrameScan(dict.fromkeys(names, I64), frame._df, None) + + +def named_col(name: str) -> expr.NamedExpr: + return expr.NamedExpr(name, expr.Col(I64, name)) + + +def make_sort( + child: IR, + *names: str, + descending: tuple[bool, ...] | None = None, + nulls_last: tuple[bool, ...] | None = None, +) -> Sort: + if descending is None: + descending = (False,) * len(names) + if nulls_last is None: + nulls_last = (False,) * len(names) + order, null_order = sort_order( + descending, nulls_last=nulls_last, num_keys=len(names) + ) + return Sort( + child.schema, + tuple(named_col(name) for name in names), + order, + null_order, + stable=False, + zlice=None, + df=child, + ) + + +def make_hint_sorted( + child: IR, + *names: str, + descending: tuple[bool, ...] | None = None, + nulls_last: tuple[bool, ...] | None = None, +) -> MapFunction: + if descending is None: + descending = (False,) * len(names) + if nulls_last is None: + nulls_last = (False,) * len(names) + return MapFunction( + child.schema, "hint_sorted", (names, descending, nulls_last), child + ) + + +def make_groupby(child: IR, *names: str, maintain_order: bool = False) -> GroupBy: + return GroupBy( + dict.fromkeys(names, I64), + tuple(named_col(name) for name in names), + (), + maintain_order=maintain_order, + zlice=None, + df=child, + ) + + +def test_sort_creates_order_partition_request() -> None: + scan = make_scan("a", "b") + sort = make_sort( + scan, + "a", + "b", + descending=(False, True), + nulls_last=(True, False), + ) + + requests = collect_partitioning_requests(sort) + + assert requests[scan] == ( + OrderPartitioningRequest( + ( + NamedOrderKey("a", descending=False, nulls_last=True), + NamedOrderKey("b", descending=True, nulls_last=False), + ) + ), + ) + + +def test_select_remaps_order_partition_request() -> None: + scan = make_scan("a", "b") + select = Select( + {"x": I64, "b": I64}, + (expr.NamedExpr("x", expr.Col(I64, "a")), named_col("b")), + should_broadcast=False, + df=scan, + ) + sort = make_sort(select, "x") + + requests = collect_partitioning_requests(sort) + + assert requests[scan] == ( + OrderPartitioningRequest( + (NamedOrderKey("a", descending=False, nulls_last=False),) + ), + ) + + +def test_non_column_sort_does_not_create_request() -> None: + scan = make_scan("a") + order, null_order = sort_order((False,), nulls_last=(False,), num_keys=1) + sort = Sort( + scan.schema, + (expr.NamedExpr("literal", expr.Literal(I64, 1)),), + order, + null_order, + stable=False, + zlice=None, + df=scan, + ) + + requests = collect_partitioning_requests(sort) + + assert requests == {} + + +def test_hint_sorted_creates_order_partition_request() -> None: + scan = make_scan("a", "b") + hint_sorted = make_hint_sorted(scan, "a", descending=(True,)) + + requests = collect_partitioning_requests(hint_sorted) + + assert requests[scan] == ( + OrderPartitioningRequest( + (NamedOrderKey("a", descending=True, nulls_last=False),) + ), + ) + + +def test_select_remaps_strict_partition_request() -> None: + scan = make_scan("a") + select = Select( + {"x": I64}, + (expr.NamedExpr("x", expr.Col(I64, "a")),), + should_broadcast=False, + df=scan, + ) + right = make_scan("x") + join = Join( + {"x": I64}, + (named_col("x"),), + (named_col("x"),), + ("Inner", False, None, "_right", True, "none"), + select, + right, + ) + + requests = collect_partitioning_requests(join) + + assert requests[scan] == (StrictPartitioningRequest(("a",)),) + assert requests[right] == (StrictPartitioningRequest(("x",)),) + + +def test_cross_join_does_not_create_strict_partition_request() -> None: + left = make_scan("a") + right = make_scan("x") + join = Join( + {"a": I64, "x": I64}, + (), + (), + ("Cross", False, None, "_right", True, "none"), + left, + right, + ) + + requests = collect_partitioning_requests(join) + + assert requests == {} + + +def test_maintain_order_groupby_does_not_create_strict_partition_request() -> None: + scan = make_scan("a") + groupby = make_groupby(scan, "a", maintain_order=True) + + requests = collect_partitioning_requests(groupby) + + assert requests == {} + + +def test_select_drops_order_request_on_non_column_output() -> None: + scan = make_scan("a") + select = Select( + {"x": I64}, + (expr.NamedExpr("x", expr.Literal(I64, 1)),), + should_broadcast=False, + df=scan, + ) + sort = make_sort(select, "x") + + requests = collect_partitioning_requests(sort) + + assert requests[select] == ( + OrderPartitioningRequest( + (NamedOrderKey("x", descending=False, nulls_last=False),) + ), + ) + assert scan not in requests + + +def test_select_drops_strict_request_on_non_column_output() -> None: + scan = make_scan("a") + select = Select( + {"x": I64}, + (expr.NamedExpr("x", expr.Literal(I64, 1)),), + should_broadcast=False, + df=scan, + ) + right = make_scan("x") + join = Join( + {"x": I64}, + (named_col("x"),), + (named_col("x"),), + ("Inner", False, None, "_right", True, "none"), + select, + right, + ) + + requests = collect_partitioning_requests(join) + + assert requests[select] == (StrictPartitioningRequest(("x",)),) + assert requests[right] == (StrictPartitioningRequest(("x",)),) + assert scan not in requests + + +def test_hint_sorted_keeps_declared_order_with_compatible_downstream_sort() -> None: + scan = make_scan("a", "b") + hint_sorted = make_hint_sorted(scan, "a") + sort = make_sort(hint_sorted, "a") + + requests = collect_partitioning_requests(sort) + + assert requests[scan] == ( + OrderPartitioningRequest( + (NamedOrderKey("a", descending=False, nulls_last=False),) + ), + ) + + +def test_hint_sorted_keeps_declared_order_with_extended_downstream_sort() -> None: + scan = make_scan("a", "b") + hint_sorted = make_hint_sorted(scan, "a") + sort = make_sort(hint_sorted, "a", "b") + + requests = collect_partitioning_requests(sort) + + assert requests[scan] == ( + OrderPartitioningRequest( + (NamedOrderKey("a", descending=False, nulls_last=False),) + ), + ) + + +def test_hint_sorted_keeps_declared_order_with_incompatible_downstream_sort() -> None: + scan = make_scan("a", "b") + hint_sorted = make_hint_sorted(scan, "a", descending=(True,)) + sort = make_sort(hint_sorted, "a") + + requests = collect_partitioning_requests(sort) + + assert requests[scan] == ( + OrderPartitioningRequest( + (NamedOrderKey("a", descending=True, nulls_last=False),) + ), + ) + + +def test_groupby_remaps_order_partition_request() -> None: + scan = make_scan("a", "b") + groupby = make_groupby( + Select( + {"key": I64}, + (expr.NamedExpr("key", expr.Col(I64, "a")),), + should_broadcast=False, + df=scan, + ), + "key", + ) + sort = make_sort(groupby, "key") + + requests = collect_partitioning_requests(sort) + + assert requests[scan] == ( + OrderPartitioningRequest( + (NamedOrderKey("a", descending=False, nulls_last=False),), + strict_key_count=1, + ), + ) + + +def test_fanout_keeps_more_specific_compatible_order_request() -> None: + scan = make_scan("a", "b") + root = Union( + scan.schema, + None, + False, # noqa: FBT003 + make_sort(scan, "a"), + make_sort(scan, "a", "b"), + ) + + requests = collect_partitioning_requests(root) + + assert requests[scan] == ( + OrderPartitioningRequest( + ( + NamedOrderKey("a", descending=False, nulls_last=False), + NamedOrderKey("b", descending=False, nulls_last=False), + ) + ), + ) + + +def test_fanout_marks_compatible_order_request_as_strict() -> None: + scan = make_scan("a", "b") + right = make_scan("a", "right_value") + join = Join( + {"a": I64, "b": I64, "right_value": I64}, + (named_col("a"),), + (named_col("a"),), + ("Inner", False, None, "_right", True, "none"), + scan, + right, + ) + root = Union( + scan.schema, + None, + False, # noqa: FBT003 + make_sort(scan, "a", "b"), + join, + ) + + requests = collect_partitioning_requests(root) + + assert requests[scan] == ( + OrderPartitioningRequest( + ( + NamedOrderKey("a", descending=False, nulls_last=False), + NamedOrderKey("b", descending=False, nulls_last=False), + ), + strict_key_count=1, + ), + ) + assert requests[right] == (StrictPartitioningRequest(("a",)),) + + +def test_fanout_merges_compatible_strict_requests() -> None: + scan = make_scan("a", "b") + root = Union( + scan.schema, + None, + False, # noqa: FBT003 + make_groupby(scan, "a"), + make_groupby(scan, "a", "b"), + ) + + requests = collect_partitioning_requests(root) + + assert requests[scan] == (StrictPartitioningRequest(("a",)),) + + +def test_fanout_keeps_incompatible_strict_candidates() -> None: + scan = make_scan("a", "b") + root = Union( + scan.schema, + None, + False, # noqa: FBT003 + make_groupby(scan, "a"), + make_groupby(scan, "b"), + ) + + requests = collect_partitioning_requests(root) + + assert set(requests[scan]) == { + StrictPartitioningRequest(("a",)), + StrictPartitioningRequest(("b",)), + } + + +def test_fanout_keeps_incompatible_order_and_strict_candidates() -> None: + scan = make_scan("a", "b") + root = Union( + scan.schema, + None, + False, # noqa: FBT003 + make_sort(scan, "a"), + make_groupby(scan, "b"), + ) + + requests = collect_partitioning_requests(root) + + assert set(requests[scan]) == { + OrderPartitioningRequest( + (NamedOrderKey("a", descending=False, nulls_last=False),) + ), + StrictPartitioningRequest(("b",)), + } + + +def test_compatible_request_merging_is_not_directional() -> None: + scan = make_scan("a", "b") + root = Union( + scan.schema, + None, + False, # noqa: FBT003 + make_groupby(scan, "a", "b"), + make_groupby(scan, "a"), + ) + assert collect_partitioning_requests(root)[scan] == ( + StrictPartitioningRequest(("a",)), + ) + + scan = make_scan("a", "b") + root = Union( + scan.schema, + None, + False, # noqa: FBT003 + make_groupby(scan, "b"), + make_sort(scan, "a"), + ) + assert set(collect_partitioning_requests(root)[scan]) == { + StrictPartitioningRequest(("b",)), + OrderPartitioningRequest( + (NamedOrderKey("a", descending=False, nulls_last=False),) + ), + } + + scan = make_scan("a", "b") + root = Union( + scan.schema, + None, + False, # noqa: FBT003 + make_sort(scan, "a", "b"), + make_sort(scan, "a"), + ) + assert collect_partitioning_requests(root)[scan] == ( + OrderPartitioningRequest( + ( + NamedOrderKey("a", descending=False, nulls_last=False), + NamedOrderKey("b", descending=False, nulls_last=False), + ) + ), + ) + + +def test_conflicting_fanout_keeps_candidate_requests() -> None: + scan = make_scan("a", "b") + root = Union( + scan.schema, + None, + False, # noqa: FBT003 + make_sort(scan, "a"), + make_sort(scan, "b"), + ) + + requests = collect_partitioning_requests(root) + + assert set(requests[scan]) == { + OrderPartitioningRequest( + (NamedOrderKey("a", descending=False, nulls_last=False),) + ), + OrderPartitioningRequest( + (NamedOrderKey("b", descending=False, nulls_last=False),) + ), + } + + +def test_repeated_fanout_candidate_is_merged() -> None: + scan = make_scan("a", "b") + root = Union( + scan.schema, + None, + False, # noqa: FBT003 + make_sort(scan, "a"), + make_sort(scan, "b"), + make_sort(scan, "a"), + ) + + requests = collect_partitioning_requests(root) + + assert len(requests[scan]) == 2 + assert set(requests[scan]) == { + OrderPartitioningRequest( + (NamedOrderKey("a", descending=False, nulls_last=False),) + ), + OrderPartitioningRequest( + (NamedOrderKey("b", descending=False, nulls_last=False),) + ), + }