diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/collectives/sort.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/collectives/sort.py index 5b5a720c095c..90f008c24352 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/collectives/sort.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/collectives/sort.py @@ -701,6 +701,18 @@ def _sort_to_order_keys(ir: Sort) -> list[OrderKey]: ] +def _can_sort_chunkwise( + ordering: Ordering | None, order_keys: Sequence[OrderKey] +) -> bool: + """Return true when ordering avoids a global sort.""" + if ordering is None: + return False + keys = tuple(ordering.keys) + return keys == tuple(order_keys[: len(keys)]) and ( + len(keys) == len(order_keys) or ordering.strict_boundaries + ) + + def _build_order_scheme( context: Context, order_keys: list[OrderKey], @@ -832,10 +844,10 @@ async def sort_actor( partitioning = NormalizedPartitioning.from_keys( metadata_in.partitioning, comm.nranks, keys=order_keys ) - if partitioning.is_ordered( - order_keys, - level="local" if metadata_in.duplicated else "flat", - ): + ordering = partitioning.get_ordering( + level="local" if metadata_in.duplicated else "flat" + ) + if _can_sort_chunkwise(ordering, order_keys): if tracer is not None: tracer.decision = "already_sorted" await chunkwise_evaluate( diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/groupby.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/groupby.py index 40b48f821997..316278253926 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/groupby.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/groupby.py @@ -891,6 +891,13 @@ async def groupby_actor( fully_partitioned = partitioning.is_strictly_partitioned( level=partitioning_level, ) + input_ordering = None if metadata_in.duplicated else partitioning.get_ordering() + can_adjust_preserving_order = ( + isinstance(ir, GroupBy) + and preserves_output_order + and input_ordering is not None + and not _has_stable_sorted_agg(ir.agg_requests) + ) fallback_case = ( # NOTE: This criteria means that we fell back # to one partition at lowering time. @@ -936,7 +943,7 @@ async def groupby_actor( ir_context, ch_in, target_partition_size, - allow_early_exit=not maintain_order, + allow_early_exit=not maintain_order or can_adjust_preserving_order, ) skip_global_comm = metadata_in.duplicated or isinstance( @@ -952,7 +959,7 @@ async def groupby_actor( collective_ids, target_partition_size, skip_global_comm, - maintain_order, + maintain_order and not can_adjust_preserving_order, tracer, ) @@ -969,11 +976,7 @@ async def groupby_actor( aggregated=aggregated, tracer=tracer, ) - elif not metadata_in.duplicated and partitioning.is_ordered( - group_keys, - level="flat", - ): - assert isinstance(partitioning.inter_rank_scheme, OrderScheme) + elif input_ordering is not None: await _ordered_adjust_reduce( context, comm, @@ -986,7 +989,7 @@ async def groupby_actor( target_partition_size, aggregated=aggregated, input_drained=input_drained, - input_ordering=partitioning.inter_rank_scheme.orderings[0], + input_ordering=input_ordering, preserves_output_order=preserves_output_order, tracer=tracer, ) diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/join.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/join.py index c2577992bfdd..d0c10e750b81 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/join.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/join.py @@ -139,6 +139,12 @@ class OrderedJoinStrategy: JoinStrategy: TypeAlias = ( BroadcastJoinStrategy | ShuffleJoinStrategy | OrderedJoinStrategy ) +OrderedJoinDecision: TypeAlias = Literal[ + "ordered_aligned", + "ordered_adjust_left", + "ordered_adjust_right", + "ordered_adjust_both", +] @define_actor() @@ -538,6 +544,17 @@ def _make_ordered_strategy( if not _ordering_prefix_matches(right_ordering, reference, right_key_indices): return None + left_output_ordering = _update_ordering_indices(reference, left_key_indices) + output_ordering = _update_ordering_indices(reference, output_key_indices) + if not ( + left_ordering.locally_ordered + and join_preserves_side_order(ir.options[5], "left") + ): + left_output_ordering = left_output_ordering.with_locally_ordered( + locally_ordered=False + ) + output_ordering = output_ordering.with_locally_ordered(locally_ordered=False) + return OrderedJoinStrategy( output_indices=output_key_indices, left_indices=left_key_indices, @@ -546,9 +563,11 @@ def _make_ordered_strategy( right_keys=right_keys[:reference_key_count], left_input_ordering=left_ordering, right_input_ordering=right_ordering, - left_output_ordering=_update_ordering_indices(reference, left_key_indices), - right_output_ordering=_update_ordering_indices(reference, right_key_indices), - output_ordering=_update_ordering_indices(reference, output_key_indices), + left_output_ordering=left_output_ordering, + right_output_ordering=_update_ordering_indices( + reference, right_key_indices + ).with_locally_ordered(locally_ordered=False), + output_ordering=output_ordering, ) @@ -744,6 +763,19 @@ def _local_count_for_ordering(comm: Communicator, ordering: Ordering) -> int: return stop - start +def _ordered_join_decision( + *, left_aligned: bool, right_aligned: bool +) -> OrderedJoinDecision: + """Return the trace decision for ordered join-side alignment.""" + if left_aligned and right_aligned: + return "ordered_aligned" + if left_aligned: + return "ordered_adjust_right" + if right_aligned: + return "ordered_adjust_left" + return "ordered_adjust_both" + + async def _adjust_ordered_join_side( context: Context, comm: Communicator, @@ -754,21 +786,30 @@ async def _adjust_ordered_join_side( input_ordering: Ordering, output_ordering: Ordering, *, + already_aligned: bool, collective_id: int, ) -> None: """Send metadata, then align one join side to output_ordering.""" - await send_metadata( - ch_out, - context, - ChannelMetadata( - local_count=_local_count_for_ordering(comm, output_ordering), - partitioning=Partitioning( - OrderScheme([output_ordering]), - local="inherit", - ), - duplicated=False, + output_metadata = ChannelMetadata( + local_count=_local_count_for_ordering(comm, output_ordering), + partitioning=Partitioning( + OrderScheme([output_ordering]), + local="inherit", ), + duplicated=False, ) + if already_aligned: + await replay_buffered_channel( + context, + ch_out, + ch_in, + (), + output_metadata, + trace_ir=schema_ir, + ) + return + + await send_metadata(ch_out, context, output_metadata) await adjust_ordering( context, comm, @@ -796,6 +837,12 @@ async def _ordered_join( tracer: ActorTracer | None, ) -> None: """Align ordered inputs to common boundaries, then join partition-wise.""" + left_boundaries_aligned = strategy.left_input_ordering.boundaries_aligned_with( + strategy.left_output_ordering, context.br() + ) + right_boundaries_aligned = strategy.right_input_ordering.boundaries_aligned_with( + strategy.right_output_ordering, context.br() + ) metadata_out = ChannelMetadata( local_count=_local_count_for_ordering(comm, strategy.output_ordering), partitioning=Partitioning( @@ -806,10 +853,26 @@ async def _ordered_join( ) await send_metadata(ch_out, context, metadata_out) - await gather_in_task_group( + left_metadata, right_metadata = await gather_in_task_group( recv_metadata(ch_left, context), recv_metadata(ch_right, context), ) + left_target_count = _local_count_for_ordering(comm, strategy.left_output_ordering) + right_target_count = _local_count_for_ordering(comm, strategy.right_output_ordering) + left_aligned = ( + left_boundaries_aligned + and not left_metadata.duplicated + and left_metadata.local_count == left_target_count + ) + right_aligned = ( + right_boundaries_aligned + and not right_metadata.duplicated + and right_metadata.local_count == right_target_count + ) + if tracer is not None: + tracer.decision = _ordered_join_decision( + left_aligned=left_aligned, right_aligned=right_aligned + ) ch_left_adjusted = context.create_channel() ch_right_adjusted = context.create_channel() async with shutdown_on_error( @@ -829,6 +892,7 @@ async def _ordered_join( ch_left, strategy.left_input_ordering, strategy.left_output_ordering, + already_aligned=left_aligned, collective_id=collective_ids.pop(0), ), _adjust_ordered_join_side( @@ -840,6 +904,7 @@ async def _ordered_join( ch_right, strategy.right_input_ordering, strategy.right_output_ordering, + already_aligned=right_aligned, collective_id=collective_ids.pop(0), ), _join_chunks( diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/utils.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/utils.py index f881c73f67db..c39f1d1c3283 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/utils.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/utils.py @@ -54,6 +54,7 @@ Callable, Coroutine, Generator, + Iterable, Iterator, Sequence, ) @@ -1278,7 +1279,7 @@ async def replay_buffered_channel( context: Context, ch_out: Channel[TableChunk], ch_in: Channel[TableChunk], - buffered_chunks: ChunkStore, + buffered_chunks: Iterable[Message], metadata: ChannelMetadata, *, trace_ir: IR, @@ -1295,7 +1296,7 @@ async def replay_buffered_channel( ch_in The buffered input channel. buffered_chunks - The buffered chunks to yield first. + Buffered messages to yield first. May be empty. metadata The metadata to send to the output channel. trace_ir @@ -1358,29 +1359,6 @@ def _scheme_is_strict(scheme: PartitioningScheme) -> bool: return scheme.orderings[0].strict_boundaries return True - @staticmethod - def _ordering_covers_keys( - ordering: Ordering, - order_keys: Sequence[int | OrderKey], - ) -> bool: - """True when an ordering covers the requested sort keys.""" - ordering_keys = ordering.keys - if len(ordering.keys) < len(order_keys): - # If we are only sorted on a subset of the keys, we need strict - # boundaries to know later keys cannot interleave across chunks. - if not ordering.strict_boundaries: - return False - order_keys = order_keys[: len(ordering.keys)] - else: - ordering_keys = ordering.keys[: len(order_keys)] - for current, target in zip(ordering_keys, order_keys, strict=True): - if isinstance(target, OrderKey): - if current != target: - return False - elif current.column_index != target: - return False - return True - def is_strictly_partitioned( self, *, @@ -1389,17 +1367,10 @@ def is_strictly_partitioned( """True if data is strictly partitioned at the requested level.""" return self._scheme_is_strict(self._scheme_for_level(level)) - def is_ordered( - self, - order_keys: Sequence[int | OrderKey], - *, - level: PartitioningLevel = "flat", - ) -> bool: - """True if the selected ordering covers order_keys.""" + def get_ordering(self, *, level: PartitioningLevel = "flat") -> Ordering | None: + """Return the normalized ordering for the requested partitioning level.""" scheme = self._scheme_for_level(level) - if not isinstance(scheme, OrderScheme): - return False - return self._ordering_covers_keys(scheme.orderings[0], order_keys) + return scheme.orderings[0] if isinstance(scheme, OrderScheme) else None def is_aligned_with( self, other: NormalizedPartitioning, br: BufferResource diff --git a/python/cudf_polars/tests/streaming/test_groupby.py b/python/cudf_polars/tests/streaming/test_groupby.py index 28fbe8086c7b..8b169abd0207 100644 --- a/python/cudf_polars/tests/streaming/test_groupby.py +++ b/python/cudf_polars/tests/streaming/test_groupby.py @@ -14,15 +14,19 @@ import polars as pl import pylibcudf as plc +from cudf_streaming.channel_metadata import OrderScheme from cudf_streaming.table_chunk import TableChunk +from cudf_polars import Translator from cudf_polars.containers import DataFrame, DataType from cudf_polars.dsl import expr from cudf_polars.dsl.ir import Distinct, Empty, GroupBy from cudf_polars.engine.options import StreamingOptions from cudf_polars.streaming.actor_graph import groupby as groupby_actor_graph from cudf_polars.streaming.actor_graph.collectives.shuffle import ShuffleManager +from cudf_polars.streaming.actor_graph.core import evaluate_logical_plan from cudf_polars.testing.asserts import assert_gpu_result_equal +from cudf_polars.utils.config import ConfigOptions @pytest.fixture(scope="module") @@ -132,6 +136,58 @@ def test_order_sensitive_execution_does_not_imply_output_order() -> None: assert not distinct.preserves_output_order +def test_groupby_adjusts_truncated_ordering_with_maintain_order( + spmd_engine_factory, +) -> None: + """GroupBy can adjust an ordered prefix without tree-reducing.""" + engine = spmd_engine_factory( + StreamingOptions( + target_partition_size=1, + max_rows_per_partition=8, + fallback_mode="raise", + raise_on_fail=True, + ) + ) + df = pl.LazyFrame( + { + "DateTime": [i * 250 for i in range(128)], + "RIC": ["a", "b", "a", "b"] * 32, + "value": range(128), + } + ) + q = ( + df.sort("DateTime") + .with_columns( + pl.col("DateTime") + .cast(pl.Datetime("ns")) + .dt.truncate("1us") + .cast(pl.Int64) + .alias("ts_bucket") + ) + .group_by("ts_bucket", "RIC", maintain_order=True) + .agg(pl.col("value").sum()) + ) + assert_gpu_result_equal(q, engine=engine, check_row_order=False) + + ir = Translator(q._ldf.visit(), engine).translate_ir() + + metadata_collector = evaluate_logical_plan( + ir, ConfigOptions.from_polars_engine(engine), collect_metadata=True + )[1] + + assert metadata_collector is not None + assert len(metadata_collector) == 1 + metadata = metadata_collector[0] + assert metadata.partitioning is not None + scheme = metadata.partitioning.inter_rank + assert isinstance(scheme, OrderScheme) + assert metadata.partitioning.local == "inherit" + (ordering,) = scheme.orderings + assert tuple(key.column_index for key in ordering.keys) == (0,) + assert ordering.strict_boundaries is True + assert ordering.locally_ordered is True + + @pytest.mark.parametrize("keys", [("key",), ("key", "key2")]) @pytest.mark.parametrize("agg", ["sum", "mean", "len", "min", "max"]) def test_dynamic_groupby_basic(df, streaming_engine, keys, agg): diff --git a/python/cudf_polars/tests/streaming/test_join.py b/python/cudf_polars/tests/streaming/test_join.py index f0f03092d07b..d17112470f3b 100644 --- a/python/cudf_polars/tests/streaming/test_join.py +++ b/python/cudf_polars/tests/streaming/test_join.py @@ -12,7 +12,11 @@ import polars as pl import pylibcudf as plc -from cudf_streaming.channel_metadata import OrderKey, OrderScheme, Ordering +from cudf_streaming.channel_metadata import ( + OrderKey, + OrderScheme, + Ordering, +) from cudf_streaming.table_chunk import TableChunk from cudf_polars import Translator @@ -20,8 +24,10 @@ from cudf_polars.dsl.ir import Cache, Join from cudf_polars.dsl.traversal import traversal from cudf_polars.engine.options import StreamingOptions +from cudf_polars.streaming.actor_graph.core import evaluate_logical_plan from cudf_polars.streaming.actor_graph.join import ( _make_ordered_strategy, + _ordered_join_decision, _use_pwise_join, ) from cudf_polars.streaming.actor_graph.utils import NormalizedPartitioning @@ -32,10 +38,13 @@ from cudf_polars.testing.asserts import assert_gpu_result_equal from cudf_polars.testing.engine_utils import warns_on_spmd from cudf_polars.utils.config import ConfigOptions, StreamingExecutor +from cudf_polars.utils.versions import POLARS_VERSION_LT_138 if TYPE_CHECKING: import concurrent.futures + from cudf_streaming.channel_metadata import ChannelMetadata + @pytest.fixture def left(): @@ -170,6 +179,46 @@ def _order_partitioning( ) +def _set_sorted_join_metadata(spmd_engine_factory) -> ChannelMetadata: + engine = spmd_engine_factory( + StreamingOptions( + max_rows_per_partition=128, + target_partition_size=1, + broadcast_limit=1, + fallback_mode="raise", + raise_on_fail=True, + ), + ) + left = pl.LazyFrame({"k": [1, 2, 3], "x": [10, 20, 30]}).set_sorted("k") + right = pl.LazyFrame({"k": [1, 2, 3], "y": [100, 200, 300]}).set_sorted("k") + q = left.join(right, on="k", how="inner") + + assert_gpu_result_equal(q, engine=engine, check_row_order=False) + ir = Translator(q._ldf.visit(), engine).translate_ir() + metadata_collector = evaluate_logical_plan( + ir, ConfigOptions.from_polars_engine(engine), collect_metadata=True + )[1] + assert metadata_collector is not None + assert len(metadata_collector) == 1 + return metadata_collector[0] + + +@pytest.mark.skipif( + POLARS_VERSION_LT_138, reason="set_sorted lowers to unsupported hint ir" +) +def test_ordered_join_after_set_sorted_inputs(spmd_engine_factory): + metadata = _set_sorted_join_metadata(spmd_engine_factory) + + assert metadata.local_count == 1 + assert metadata.partitioning is not None + assert isinstance(metadata.partitioning.inter_rank, OrderScheme) + assert metadata.partitioning.local == "inherit" + (ordering,) = metadata.partitioning.inter_rank.orderings + assert tuple(key.column_index for key in ordering.keys) == (0,) + assert ordering.strict_boundaries is True + assert ordering.locally_ordered is False + + def test_ordered_join_strategy_matches_exact_ordering_width(spmd_engine): left = pl.LazyFrame({"k0": [1], "k1": [2], "x": [3]}) right = pl.LazyFrame({"k0": [1], "k1": [2], "y": [4]}) @@ -196,6 +245,25 @@ def test_ordered_join_strategy_checks_order_key_semantics(spmd_engine): assert _make_ordered_strategy(join_ir, left_partitioning, mismatched_right) is None +@pytest.mark.parametrize( + "left_aligned,right_aligned,expected", + [ + (True, True, "ordered_aligned"), + (True, False, "ordered_adjust_right"), + (False, True, "ordered_adjust_left"), + (False, False, "ordered_adjust_both"), + ], +) +def test_ordered_join_decision(left_aligned, right_aligned, expected): + assert ( + _ordered_join_decision( + left_aligned=left_aligned, + right_aligned=right_aligned, + ) + == expected + ) + + @pytest.mark.parametrize("how", ["right", "full"]) def test_ordered_join_strategy_rejects_ambiguous_output_key_metadata(spmd_engine, how): left = pl.LazyFrame({"k0": [1], "x": [2]}) diff --git a/python/cudf_polars/tests/streaming/test_metadata.py b/python/cudf_polars/tests/streaming/test_metadata.py index 7a1045760d17..e6c6b2478586 100644 --- a/python/cudf_polars/tests/streaming/test_metadata.py +++ b/python/cudf_polars/tests/streaming/test_metadata.py @@ -37,6 +37,7 @@ ) from cudf_polars.engine.options import StreamingOptions from cudf_polars.streaming.actor_graph.collectives.sort import ( + _can_sort_chunkwise, _sort_to_order_keys, ) from cudf_polars.streaming.actor_graph.core import evaluate_logical_plan @@ -1096,7 +1097,7 @@ def test_sort_output_metadata(spmd_engine_factory, by, descending, nulls_last) - @pytest.mark.parametrize( - "scheme_key_count,strict_boundaries,locally_ordered,expected", + "scheme_key_count,strict_boundaries,locally_ordered,expected_sort_compatible", [ (1, True, True, True), # prefix match + strict boundaries (1, True, False, True), # local order is not required for partitioning @@ -1106,8 +1107,12 @@ def test_sort_output_metadata(spmd_engine_factory, by, descending, nulls_last) - (2, False, True, True), # exact match + non-strict boundaries ], ) -def test_is_ordered( - spmd_engine, scheme_key_count, strict_boundaries, locally_ordered, expected +def test_get_ordering( + spmd_engine, + scheme_key_count, + strict_boundaries, + locally_ordered, + expected_sort_compatible, ) -> None: df_lf = pl.LazyFrame({"x": list(range(5)), "y": list(range(5))}) base_ir = Translator(df_lf._ldf.visit(), spmd_engine).translate_ir() @@ -1141,14 +1146,24 @@ def test_is_ordered( partitioning = NormalizedPartitioning.from_keys( meta.partitioning, nranks=1, keys=order_keys ) - assert partitioning.is_ordered(order_keys) is expected - assert partitioning.is_ordered(order_keys, level="flat") is expected - assert not partitioning.is_ordered(order_keys, level="local") + ordering = partitioning.get_ordering() + assert ordering is not None + assert tuple(k.column_index for k in ordering.keys) == tuple( + range(scheme_key_count) + ) + assert ordering.strict_boundaries is strict_boundaries + assert ordering.locally_ordered is locally_ordered + assert partitioning.get_ordering(level="flat") is not None + assert partitioning.get_ordering(level="local") is None + assert _can_sort_chunkwise(ordering, order_keys) is expected_sort_compatible nested = NormalizedPartitioning(scheme, scheme) - assert not nested.is_ordered(order_keys) - assert nested.is_ordered(order_keys, level="inter_rank") is expected - assert nested.is_ordered(order_keys, level="local") is expected + assert nested.get_ordering() is None + for level in ("inter_rank", "local"): + nested_ordering = nested.get_ordering(level=level) + assert nested_ordering is not None + sort_compatible = _can_sort_chunkwise(nested_ordering, order_keys) + assert sort_compatible is expected_sort_compatible if scheme_key_count == 2: desc, after = plc.types.Order.DESCENDING, plc.types.NullOrder.AFTER @@ -1158,6 +1173,15 @@ def test_is_ordered( (OrderKey(0, asc, after), OrderKey(1, asc, before)), ] for mismatched_keys in mismatched_order_keys: - assert not partitioning.is_ordered(mismatched_keys) - assert not nested.is_ordered(mismatched_keys, level="inter_rank") - assert not nested.is_ordered(mismatched_keys, level="local") + assert _can_sort_chunkwise(ordering, mismatched_keys) is False + mismatched = NormalizedPartitioning.from_keys( + meta.partitioning, nranks=1, keys=mismatched_keys + ) + assert mismatched.get_ordering() is None + mismatched_nested = NormalizedPartitioning.from_keys( + Partitioning(inter_rank=scheme, local=scheme), + nranks=1, + keys=mismatched_keys, + ) + assert mismatched_nested.get_ordering(level="inter_rank") is None + assert mismatched_nested.get_ordering(level="local") is None