diff --git a/docs/cudf/source/cudf_polars/api.md b/docs/cudf/source/cudf_polars/api.md index 6acff73d4f03..3c2b0d88dc32 100644 --- a/docs/cudf/source/cudf_polars/api.md +++ b/docs/cudf/source/cudf_polars/api.md @@ -67,6 +67,7 @@ Most users interact with them through `StreamingOptions` fields rather than dire .. automodule:: cudf_polars.utils.config :members: DynamicPlanningOptions, + JoinDomainPrefilterOptions, MemoryResourceConfig, ParquetOptions, StreamingExecutor, diff --git a/docs/cudf/source/cudf_polars/options.md b/docs/cudf/source/cudf_polars/options.md index 5d813e74bbe1..94ca7d52e1bf 100644 --- a/docs/cudf/source/cudf_polars/options.md +++ b/docs/cudf/source/cudf_polars/options.md @@ -108,6 +108,7 @@ Environment variables follow these patterns: | `broadcast_limit` | Maximum number of bytes for broadcast joins. | auto | | `target_partition_size` | Target partition size in bytes. Used for IO and dynamic planning. `0` means auto. | auto | | `dynamic_planning` | Dynamic planning configuration, dict or {class}`~cudf_polars.utils.config.DynamicPlanningOptions`. `None` disables. | enabled | +| `join_domain_prefilter` | Join-domain prefilter configuration, dict or {class}`~cudf_polars.utils.config.JoinDomainPrefilterOptions`. `None` disables. | enabled | | `sink_to_directory` | Whether `.sink_*()` writes its output as a directory. The `spmd`, `ray`, and `dask` engines always use `True`; passing `False` raises `ValueError`. | `True` | ### Category: `engine` diff --git a/python/cudf_polars/cudf_polars/engine/options.py b/python/cudf_polars/cudf_polars/engine/options.py index 559c05be8c94..1ad8bdf4ee13 100644 --- a/python/cudf_polars/cudf_polars/engine/options.py +++ b/python/cudf_polars/cudf_polars/engine/options.py @@ -25,6 +25,7 @@ from cudf_polars.utils.config import ( DynamicPlanningOptions, + JoinDomainPrefilterOptions, ParquetOptions, ) @@ -247,6 +248,14 @@ class StreamingOptions: Env: ``CUDF_POLARS__EXECUTOR__DYNAMIC_PLANNING``. Default: enabled. Category: executor. + join_domain_prefilter + Join-domain prefilter config, dict or + :class:`~cudf_polars.utils.config.JoinDomainPrefilterOptions`. ``None`` + disables the rewrite. + Env: ``CUDF_POLARS__EXECUTOR__JOIN_DOMAIN_PREFILTER`` and + ``CUDF_POLARS__EXECUTOR__JOIN_DOMAIN_PREFILTER__*``. + Default: enabled. + Category: executor. sink_to_directory Whether multi-partition sink operations should write to a directory rather than a single file. The ``spmd``/``ray``/``dask`` engines @@ -341,6 +350,9 @@ class StreamingOptions: dynamic_planning: dict[str, Any] | DynamicPlanningOptions | None | Unspecified = ( _opt("executor") ) + join_domain_prefilter: ( + dict[str, Any] | JoinDomainPrefilterOptions | None | Unspecified + ) = _opt("executor") sink_to_directory: bool | Unspecified = _opt( "executor", "CUDF_POLARS__EXECUTOR__SINK_TO_DIRECTORY", parse_boolean ) diff --git a/python/cudf_polars/cudf_polars/streaming/join_domain_prefilter.py b/python/cudf_polars/cudf_polars/streaming/join_domain_prefilter.py new file mode 100644 index 000000000000..c9e22f7003f2 --- /dev/null +++ b/python/cudf_polars/cudf_polars/streaming/join_domain_prefilter.py @@ -0,0 +1,947 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Generic derived key-domain prefilters for streaming joins.""" + +from __future__ import annotations + +from dataclasses import dataclass +from functools import singledispatch +from typing import TYPE_CHECKING, Any, Literal, TypedDict + +from cudf_polars.dsl import expr +from cudf_polars.dsl.ir import ( + IR, + Cache, + DataFrameScan, + Distinct, + Filter, + GroupBy, + HStack, + Join, + Projection, + Scan, + Select, + Sort, +) +from cudf_polars.dsl.tracing import Scope, log +from cudf_polars.dsl.traversal import ( + CachingVisitor, + post_traversal, + reuse_if_unchanged, + traversal, +) +from cudf_polars.dsl.utils.replace import replace + +if TYPE_CHECKING: + from collections.abc import Iterable, Sequence + + from cudf_polars.streaming.base import StatsCollector + from cudf_polars.typing import GenericTransformer + from cudf_polars.utils.config import ConfigOptions, StreamingExecutor + + +@dataclass(frozen=True) +class _Producer: + """A subtree and its bound column names at an insertion point.""" + + node: IR + columns: tuple[str, ...] + rows: int + cost: int + + @property + def column(self) -> str: + """First bound column in the producer.""" + return self.columns[0] + + +@dataclass(frozen=True) +class _Candidate: + """A derived key-domain prefilter candidate.""" + + mode: Literal["simple", "composite", "reused"] + target_side: Literal["left", "right"] + target: _Producer + target_key: expr.Col + domain: _Producer + domain_key: expr.Col + target_rows: int + constraint_domain: _Producer | None = None + domain_constraint_key: expr.Col | None = None + target_constraint_key: expr.Col | None = None + + @property + def domain_rows(self) -> int: + """Estimated rows in the domain input.""" + return self.domain.rows + + @property + def domain_cost(self) -> int: + """Estimated scan-row cost to build the domain input.""" + return self.domain.cost + + @property + def score(self) -> tuple[int, int, int]: + """Prefer composite filters, then cheaper constraint/domain inputs.""" + mode_score = {"composite": 0, "simple": 1, "reused": 2}[self.mode] + constraint_cost = ( + self.constraint_domain.cost + if self.constraint_domain is not None + else self.domain.cost + ) + return ( + mode_score, + constraint_cost, + self.domain.cost, + ) + + +class _RewriteState(TypedDict): + """State shared by the join-domain prefilter DAG rewrite.""" + + threshold: float + trace: bool + stats: StatsCollector + row_estimates: dict[IR, int | None] + source_costs: dict[IR, int | None] + source_counts: dict[IR, int] + selective_nodes: set[IR] + + +def optimize_join_domain_prefilters( + ir: IR, + stats: StatsCollector, + config_options: ConfigOptions[StreamingExecutor], +) -> IR: + """ + Insert generic semi-join key-domain prefilters before streaming lowering. + + The rewrite is intentionally conservative: only inner joins with simple + column equality keys are considered, and the original full join remains + after every inserted row-reduction semi join. + """ + options = config_options.executor.join_domain_prefilter + if options is None: + return ir + threshold = options.threshold + trace = options.trace + if threshold == 0: + return ir + + row_estimates = _estimate_row_counts(ir, stats) + source_costs, source_counts = _estimate_source_stats(ir, row_estimates) + state = _RewriteState( + threshold=threshold, + trace=trace, + stats=stats, + row_estimates=row_estimates, + source_costs=source_costs, + source_counts=source_counts, + selective_nodes=_collect_selective_nodes(ir), + ) + mapper: GenericTransformer[IR, IR, _RewriteState] = CachingVisitor( + _rewrite, state=state + ) + return mapper(ir) + + +@singledispatch +def _rewrite(node: IR, rec: GenericTransformer[IR, IR, _RewriteState]) -> IR: + raise AssertionError + + +@_rewrite.register(IR) +def _(node: IR, rec: GenericTransformer[IR, IR, _RewriteState]) -> IR: + return reuse_if_unchanged(node, rec) + + +@_rewrite.register(Join) +def _(node: Join, rec: GenericTransformer[IR, IR, _RewriteState]) -> IR: + original = node + rewritten = reuse_if_unchanged(node, rec) + assert isinstance(rewritten, Join) + node = rewritten + if node is original: + row_estimates = rec.state["row_estimates"] + source_costs = rec.state["source_costs"] + source_counts = rec.state["source_counts"] + selective_nodes = rec.state["selective_nodes"] + else: + # Child rewrites introduce new semi joins and reconstructed ancestors. + # Re-analyze that current subtree so parent joins can use the derived + # selectivity and cardinality when ranking their own candidates. + row_estimates = _estimate_row_counts(node, rec.state["stats"]) + source_costs, source_counts = _estimate_source_stats(node, row_estimates) + selective_nodes = _collect_selective_nodes(node) + candidate, reason = _select_candidate( + node, + rec.state["threshold"], + row_estimates, + source_costs, + source_counts, + selective_nodes, + ) + if rec.state["trace"]: + _trace_decision(node, rec.state["threshold"], candidate, reason, source_costs) + if candidate is None: + return node + + left, right = node.children + domain = _make_domain(candidate, node) + target = candidate.target + target_filter = _make_semi_join( + target.node, + expr.Col(target.node.schema[target.column], target.column), + domain, + expr.Col(domain.schema[candidate.domain_key.name], candidate.domain_key.name), + nulls_equal=node.options[1], + suffix=node.options[3], + ) + # A DAG may share the target with the domain side, so only rewrite the + # side for which this candidate was selected. + if candidate.target_side == "left": + (left,) = replace([left], {candidate.target.node: target_filter}) + else: + (right,) = replace([right], {candidate.target.node: target_filter}) + return node.reconstruct((left, right)) + + +def _select_candidate( + ir: Join, + threshold: float, + row_estimates: dict[IR, int | None], + source_costs: dict[IR, int | None], + source_counts: dict[IR, int], + selective_nodes: set[IR], +) -> tuple[_Candidate | None, str]: + if ir.options[0] != "Inner": + return None, "not_inner_join" + if ir.options[2] is not None: + return None, "sliced_join" + if ir.options[5] != "none": + return None, "maintain_order" + + left_keys = _simple_keys(ir.left_on) + right_keys = _simple_keys(ir.right_on) + if len(left_keys) != len(ir.left_on) or len(right_keys) != len(ir.right_on): + return None, "non_column_join_key" + + candidates: list[_Candidate] = [] + left: tuple[Literal["left", "right"], IR, tuple[expr.Col, ...]] = ( + "left", + ir.children[0], + left_keys, + ) + right: tuple[Literal["left", "right"], IR, tuple[expr.Col, ...]] = ( + "right", + ir.children[1], + right_keys, + ) + for (target_side, target_child, target_keys), ( + _, + domain_child, + domain_keys, + ) in ((left, right), (right, left)): + candidates.extend( + _composite_candidates( + target_side, + target_child, + domain_child, + target_keys, + domain_keys, + threshold, + row_estimates, + source_costs, + selective_nodes, + ) + ) + candidates.extend( + _simple_candidates( + target_side, + target_child, + domain_child, + target_keys, + domain_keys, + threshold, + row_estimates, + source_costs, + source_counts, + selective_nodes, + ) + ) + candidates.extend( + _reused_semi_domain_candidates( + target_side, + target_child, + domain_child, + target_keys, + domain_keys, + row_estimates, + source_costs, + source_counts, + selective_nodes, + ) + ) + + if not candidates: + return None, "no_profitable_domain" + return min(candidates, key=lambda c: c.score), "applied" + + +def _simple_keys(keys: Sequence[expr.NamedExpr]) -> tuple[expr.Col, ...]: + return tuple(key.value for key in keys if isinstance(key.value, expr.Col)) + + +def _simple_candidates( + target_side: Literal["left", "right"], + target_child: IR, + domain_child: IR, + target_keys: tuple[expr.Col, ...], + domain_keys: tuple[expr.Col, ...], + threshold: float, + row_estimates: dict[IR, int | None], + source_costs: dict[IR, int | None], + source_counts: dict[IR, int], + selective_nodes: set[IR], +) -> Iterable[_Candidate]: + for target_key, domain_key in zip(target_keys, domain_keys, strict=True): + target = _largest_key_source( + target_child, target_key.name, row_estimates, source_costs + ) + if target is None: + continue + domain = _smallest_key_producer( + domain_child, + domain_key.name, + row_estimates, + source_costs, + selective_nodes, + require_selective=True, + ) + if domain is None: + continue + if domain.rows / target.rows > threshold: + continue + if _contains_identity(target.node, domain.node): + continue + if source_counts.get(domain.node) == 1 and _has_filtering_semi_ancestor( + target_child, target.node + ): + continue + if not _domain_cost_is_small( + domain.node, target.node, threshold, row_estimates, source_costs + ): + continue + yield _Candidate( + mode="simple", + target_side=target_side, + target=target, + target_key=target_key, + domain=domain, + domain_key=domain_key, + target_rows=target.rows, + ) + + +def _reused_semi_domain_candidates( + target_side: Literal["left", "right"], + target_child: IR, + domain_child: IR, + target_keys: tuple[expr.Col, ...], + domain_keys: tuple[expr.Col, ...], + row_estimates: dict[IR, int | None], + source_costs: dict[IR, int | None], + source_counts: dict[IR, int], + selective_nodes: set[IR], +) -> Iterable[_Candidate]: + for target_key, domain_key in zip(target_keys, domain_keys, strict=True): + for domain, semi_domain_key in _filtering_semi_domains( + domain_child, + domain_key, + row_estimates, + source_costs, + selective_nodes, + ): + target = _largest_key_reusable_target( + target_child, target_key.name, domain.node, row_estimates, source_costs + ) + if target is None: + continue + if domain.rows > target.rows: + continue + if _contains_identity(target.node, domain.node): + continue + if source_counts.get(domain.node) == 1 and _has_filtering_semi_ancestor( + target_child, target.node + ): + continue + yield _Candidate( + mode="reused", + target_side=target_side, + target=target, + target_key=target_key, + domain=domain, + domain_key=semi_domain_key, + target_rows=target.rows, + ) + + +def _composite_candidates( + target_side: Literal["left", "right"], + target_child: IR, + domain_child: IR, + target_keys: tuple[expr.Col, ...], + domain_keys: tuple[expr.Col, ...], + threshold: float, + row_estimates: dict[IR, int | None], + source_costs: dict[IR, int | None], + selective_nodes: set[IR], +) -> Iterable[_Candidate]: + if len(target_keys) < 2: + return + + for filter_index, (target_key, domain_key) in enumerate( + zip(target_keys, domain_keys, strict=True) + ): + target = _largest_key_source( + target_child, target_key.name, row_estimates, source_costs + ) + if target is None: + continue + + for constraint_index, ( + target_constraint_key, + domain_constraint_key, + ) in enumerate(zip(target_keys, domain_keys, strict=True)): + if constraint_index == filter_index: + continue + domain = _smallest_node_containing_all( + domain_child, + (domain_key.name, domain_constraint_key.name), + row_estimates, + source_costs, + ) + if domain is None: + continue + if domain.rows / target.rows > threshold: + continue + constraint_domain = _smallest_key_producer( + target_child, + target_constraint_key.name, + row_estimates, + source_costs, + selective_nodes, + require_selective=True, + exclude=target.node, + ) + if constraint_domain is None: + continue + if constraint_domain.rows / domain.rows > threshold: + continue + if _contains_identity(target.node, domain.node) or _contains_identity( + target.node, constraint_domain.node + ): + continue + if not _domain_cost_is_small( + domain.node, target.node, threshold, row_estimates, source_costs + ): + continue + if not _domain_cost_is_small( + constraint_domain.node, + domain.node, + threshold, + row_estimates, + source_costs, + ): + continue + yield _Candidate( + mode="composite", + target_side=target_side, + target=target, + target_key=target_key, + domain=domain, + domain_key=domain_key, + target_rows=target.rows, + constraint_domain=constraint_domain, + domain_constraint_key=domain_constraint_key, + target_constraint_key=target_constraint_key, + ) + + +def _make_domain(candidate: _Candidate, ir: Join) -> IR: + if candidate.mode in ("simple", "reused"): + return _project_bound_key( + candidate.domain.node, + candidate.domain.column, + candidate.domain_key, + ) + + assert candidate.constraint_domain is not None + assert candidate.domain_constraint_key is not None + assert candidate.target_constraint_key is not None + + constraint_domain = _project_bound_key( + candidate.constraint_domain.node, + candidate.constraint_domain.column, + candidate.target_constraint_key, + ) + constrained = _make_semi_join( + candidate.domain.node, + expr.Col( + candidate.domain.node.schema[candidate.domain.columns[1]], + candidate.domain.columns[1], + ), + constraint_domain, + expr.Col( + constraint_domain.schema[candidate.target_constraint_key.name], + candidate.target_constraint_key.name, + ), + nulls_equal=ir.options[1], + suffix=ir.options[3], + ) + return _project_bound_key( + constrained, candidate.domain.column, candidate.domain_key + ) + + +def _project_bound_key(source: IR, bound_column: str, output_key: expr.Col) -> Select: + """Project a bound source column under its join-visible key name.""" + dtype = source.schema[bound_column] + assert dtype == output_key.dtype + return Select( + {output_key.name: dtype}, + (expr.NamedExpr(output_key.name, expr.Col(dtype, bound_column)),), + True, # noqa: FBT003 + source, + ) + + +def _make_semi_join( + target: IR, + target_key: expr.Col, + domain: IR, + domain_key: expr.Col, + *, + nulls_equal: bool, + suffix: str, +) -> Join: + return Join( + target.schema, + (expr.NamedExpr(target_key.name, target_key),), + (expr.NamedExpr(domain_key.name, domain_key),), + ("Semi", nulls_equal, None, suffix, False, "none"), + target, + domain, + ) + + +def _smallest_key_producer( + root: IR, + column: str, + row_estimates: dict[IR, int | None], + source_costs: dict[IR, int | None], + selective_nodes: set[IR], + *, + require_selective: bool, + exclude: IR | None = None, +) -> _Producer | None: + candidates = [] + for node, bound_column in _column_bindings(root, column): + if node is exclude: + continue + rows = row_estimates.get(node) + if rows is None or rows <= 0: + continue + if require_selective and node not in selective_nodes: + continue + cost = source_costs.get(node) + if cost is None: + continue + producer = _Producer(node, (bound_column,), rows, cost) + candidates.append((cost, rows, len(node.schema), producer)) + if not candidates: + return None + return min(candidates, key=lambda item: (item[0], item[1], item[2]))[3] + + +def _filtering_semi_domains( + root: IR, + filtered_key: expr.Col, + row_estimates: dict[IR, int | None], + source_costs: dict[IR, int | None], + selective_nodes: set[IR], +) -> Iterable[tuple[_Producer, expr.Col]]: + """Yield domains from semi joins on the exact filtered-key lineage.""" + for node, bound_key in _column_bindings(root, filtered_key.name): + if not isinstance(node, Join) or node.options[0] != "Semi": + continue + left_keys = _simple_keys(node.left_on) + right_keys = _simple_keys(node.right_on) + if len(left_keys) != len(node.left_on) or len(right_keys) != len(node.right_on): + continue + assert len(left_keys) == len(right_keys) + for left_key, right_key in zip(left_keys, right_keys, strict=True): + if left_key.name != bound_key: + continue + domain = _smallest_key_producer( + node.children[1], + right_key.name, + row_estimates, + source_costs, + selective_nodes, + require_selective=True, + ) + if domain is not None: + yield domain, right_key + + +def _smallest_node_containing_all( + root: IR, + columns: Sequence[str], + row_estimates: dict[IR, int | None], + source_costs: dict[IR, int | None], +) -> _Producer | None: + candidates = [] + lineages = [tuple(_column_bindings(root, column)) for column in columns] + if not lineages or any(not lineage for lineage in lineages): + return None + for node, first_column in lineages[0]: + bound_columns = [first_column] + for lineage in lineages[1:]: + match = next( + ( + bound_column + for candidate, bound_column in lineage + if candidate is node + ), + None, + ) + if match is None: + break + bound_columns.append(match) + else: + rows = row_estimates.get(node) + if rows is None or rows <= 0: + continue + cost = source_costs.get(node) + if cost is None: + continue + producer = _Producer(node, tuple(bound_columns), rows, cost) + candidates.append((cost, rows, len(node.schema), producer)) + if not candidates: + return None + return min(candidates, key=lambda item: (item[0], item[1], item[2]))[3] + + +def _largest_key_source( + root: IR, + column: str, + row_estimates: dict[IR, int | None], + source_costs: dict[IR, int | None], +) -> _Producer | None: + source_candidates = [] + fallback_candidates = [] + for node, bound_column in _column_bindings(root, column): + rows = row_estimates.get(node) + if rows is None or rows <= 0: + continue + cost = source_costs.get(node) + if cost is None: + continue + item = (rows, len(node.schema), _Producer(node, (bound_column,), rows, cost)) + if isinstance(node, (Scan, DataFrameScan)): + source_candidates.append(item) + else: + fallback_candidates.append(item) + candidates = source_candidates or fallback_candidates + if not candidates: + return None + return max(candidates, key=lambda item: (item[0], -item[1]))[2] + + +def _largest_key_reusable_target( + root: IR, + column: str, + domain: IR, + row_estimates: dict[IR, int | None], + source_costs: dict[IR, int | None], +) -> _Producer | None: + """Select a target without duplicating a source used by the domain.""" + shared_cache_candidates = [] + source_candidates = [] + fallback_candidates = [] + domain_sources = _scan_sources(domain) + for order, (node, bound_column) in enumerate(_column_bindings(root, column)): + rows = row_estimates.get(node) + if rows is None or rows <= 0: + continue + cost = source_costs.get(node) + if cost is None: + continue + producer = _Producer(node, (bound_column,), rows, cost) + item = (rows, len(node.schema), -order, producer) + if _contains_identity(domain, node): + if isinstance(node, Cache): + shared_cache_candidates.append(item) + continue + if not _scan_sources(node).isdisjoint(domain_sources): + continue + if isinstance(node, (Scan, DataFrameScan)): + source_candidates.append(item) + else: + fallback_candidates.append(item) + candidates = shared_cache_candidates or source_candidates or fallback_candidates + if not candidates: + return None + return max(candidates, key=lambda item: (item[0], -item[1], item[2]))[3] + + +def _scan_sources(root: IR) -> frozenset[IR]: + """Return physical scan nodes reachable from a subtree.""" + return frozenset( + node for node in traversal([root]) if isinstance(node, (Scan, DataFrameScan)) + ) + + +def _column_bindings(root: IR, column: str) -> Iterable[tuple[IR, str]]: + """Yield exact output-to-input bindings for a column through a subplan.""" + node = root + while column in node.schema: + yield node, column + binding = _input_binding(node, column) + if binding is None: + return + node, column = binding + + +def _input_binding(node: IR, column: str) -> tuple[IR, str] | None: + """Return a proven direct input binding, stopping at ambiguous operations.""" + child = node.children[0] if len(node.children) == 1 else None + if isinstance(node, Select): + selected = next((item for item in node.exprs if item.name == column), None) + return _column_expression_binding(child, selected) + if isinstance(node, HStack): + stacked = next((item for item in node.columns if item.name == column), None) + if stacked is not None: + return _column_expression_binding(child, stacked) + return _passthrough_binding(child, column) + if isinstance(node, GroupBy): + if node.zlice is not None: + return None + key = next((item for item in node.keys if item.name == column), None) + return _column_expression_binding(child, key) + if isinstance(node, Join): + return _join_input_binding(node, column) + if isinstance(node, Distinct): + return None if node.zlice is not None else _passthrough_binding(child, column) + if isinstance(node, Sort): + return None if node.zlice is not None else _passthrough_binding(child, column) + if isinstance(node, (Cache, Filter, Projection)): + return _passthrough_binding(child, column) + return None + + +def _column_expression_binding( + child: IR | None, expression: expr.NamedExpr | None +) -> tuple[IR, str] | None: + if ( + child is not None + and expression is not None + and isinstance(expression.value, expr.Col) + and expression.value.name in child.schema + ): + return child, expression.value.name + return None + + +def _passthrough_binding(child: IR | None, column: str) -> tuple[IR, str] | None: + if child is not None and column in child.schema: + return child, column + return None + + +def _join_input_binding(node: Join, column: str) -> tuple[IR, str] | None: + if node.options[2] is not None: + return None + left, right = node.children + if node.options[0] in ("Semi", "Anti"): + return _passthrough_binding(left, column) + if node.options[0] != "Inner": + return None + bindings = [] + if column in left.schema: + bindings.append((left, column)) + suffix = node.options[3] + for right_column in right.schema: + output_column = ( + f"{right_column}{suffix}" if right_column in left.schema else right_column + ) + if output_column == column and output_column in node.schema: + bindings.append((right, right_column)) + if len(bindings) == 1: + return bindings[0] + return None + + +def _estimate_row_counts(ir: IR, stats: StatsCollector) -> dict[IR, int | None]: + estimates: dict[IR, int | None] = {} + for node in post_traversal([ir]): + if isinstance(node, (Scan, DataFrameScan)): + source = stats.scan_stats.get(node) + rows = None if source is None else source.row_count + if rows is None and isinstance(node, DataFrameScan): + rows = node.df.shape()[0] + elif isinstance( + node, (Select, Projection, HStack, Cache, Filter, Distinct, GroupBy) + ): + rows = estimates[node.children[0]] + elif isinstance(node, Join): + rows = _estimate_join_rows( + node.options[0], + estimates[node.children[0]], + estimates[node.children[1]], + ) + else: + child_estimates = [ + estimate + for child in node.children + if (estimate := estimates[child]) is not None + ] + rows = max(child_estimates, default=None) + estimates[node] = rows + return estimates + + +def _estimate_join_rows( + how: str, left_rows: int | None, right_rows: int | None +) -> int | None: + if left_rows is None: + return right_rows + if right_rows is None: + return left_rows + if how in ("Inner", "Semi", "Anti"): + return min(left_rows, right_rows) + if how == "Left": + return left_rows + if how == "Right": + return right_rows + if how == "Full": + return max(left_rows, right_rows) + return None + + +def _estimate_source_stats( + ir: IR, row_estimates: dict[IR, int | None] +) -> tuple[dict[IR, int | None], dict[IR, int]]: + source_costs: dict[IR, int | None] = {} + source_counts: dict[IR, int] = {} + source_nodes: dict[IR, set[IR]] = {} + for node in post_traversal([ir]): + sources: set[IR] + if isinstance(node, (Scan, DataFrameScan)): + sources = {node} + else: + sources = set() + for child in node.children: + sources.update(source_nodes[child]) + source_nodes[node] = sources + source_counts[node] = len(sources) + rows = [ + rows + for source in sources + if (rows := row_estimates.get(source)) is not None and rows > 0 + ] + source_costs[node] = sum(rows) if rows else row_estimates.get(node) + return source_costs, source_counts + + +def _domain_cost_is_small( + domain: IR, + target: IR, + threshold: float, + row_estimates: dict[IR, int | None], + source_costs: dict[IR, int | None], +) -> bool: + """Return whether building a domain is cheap enough for the target it reduces.""" + domain_cost = source_costs.get(domain) + target_rows = row_estimates.get(target) + if domain_cost is None or target_rows is None or target_rows <= 0: + return False + return domain_cost / target_rows <= threshold + + +def _has_filtering_semi_ancestor(root: IR, target: IR) -> bool: + """Return whether target is already below a semi join's filtered side.""" + if root is target: + return False + + for index, child in enumerate(root.children): + if not _contains_identity(child, target): + continue + if isinstance(root, Join) and root.options[0] == "Semi" and index == 0: + return True + if _has_filtering_semi_ancestor(child, target): + return True + return False + + +def _collect_selective_nodes(ir: IR) -> set[IR]: + selective: set[IR] = set() + for node in post_traversal([ir]): + if ( + (isinstance(node, Scan) and node.predicate is not None) + or isinstance(node, Filter) + or any(child in selective for child in node.children) + ): + selective.add(node) + return selective + + +def _contains_identity(root: IR, needle: IR) -> bool: + return any(node is needle for node in traversal([root])) + + +def _trace_decision( + ir: Join, + threshold: float, + candidate: _Candidate | None, + reason: str, + source_costs: dict[IR, int | None], +) -> None: + join_domain_prefilter: dict[str, Any] = { + "considered": True, + "threshold": threshold, + "reason": reason, + } + record = { + "scope": Scope.PLAN.value, + "join_domain_prefilter": join_domain_prefilter, + "actor_ir_id": ir.get_stable_id(), + "actor_ir_type": type(ir).__name__, + } + if candidate is not None: + join_domain_prefilter.update( + { + "mode": candidate.mode, + "target_side": candidate.target_side, + "target_key": candidate.target_key.name, + "domain_key": candidate.domain_key.name, + "estimated_target_rows": candidate.target_rows, + "estimated_domain_rows": candidate.domain_rows, + "estimated_target_cost": source_costs.get(candidate.target.node), + "estimated_domain_cost": candidate.domain_cost, + "target_node_type": type(candidate.target.node).__name__, + "domain_node_type": type(candidate.domain.node).__name__, + } + ) + if candidate.constraint_domain is not None: + join_domain_prefilter.update( + { + "constraint_key": candidate.target_constraint_key.name + if candidate.target_constraint_key is not None + else None, + "estimated_constraint_rows": candidate.constraint_domain.rows, + "estimated_constraint_cost": candidate.constraint_domain.cost, + } + ) + log("Join Domain Prefilter", **record) diff --git a/python/cudf_polars/cudf_polars/streaming/parallel.py b/python/cudf_polars/cudf_polars/streaming/parallel.py index 6f8734fd17b4..391a99928547 100644 --- a/python/cudf_polars/cudf_polars/streaming/parallel.py +++ b/python/cudf_polars/cudf_polars/streaming/parallel.py @@ -1,4 +1,4 @@ -# SPDX-FileCopyrightText: Copyright (c) 2024-2026, NVIDIA CORPORATION & AFFILIATES. +# SPDX-FileCopyrightText: Copyright (c) 2024-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 """Multi-partition evaluation.""" @@ -104,6 +104,12 @@ def lower_ir_graph( -------- lower_ir_node """ + from cudf_polars.streaming.join_domain_prefilter import ( + optimize_join_domain_prefilters, + ) + + ir = optimize_join_domain_prefilters(ir, stats, config_options) + state: State = { "config_options": config_options, "stats": stats, diff --git a/python/cudf_polars/cudf_polars/utils/config.py b/python/cudf_polars/cudf_polars/utils/config.py index 35100c5c38f2..f15b770cb9d9 100644 --- a/python/cudf_polars/cudf_polars/utils/config.py +++ b/python/cudf_polars/cudf_polars/utils/config.py @@ -51,6 +51,7 @@ "DaskContext", "DynamicPlanningOptions", "InMemoryExecutor", + "JoinDomainPrefilterOptions", "ParquetOptions", "RayContext", "SPMDContext", @@ -380,6 +381,52 @@ def __post_init__(self) -> None: # noqa: D105 raise TypeError("join_prefilter_trace must be a bool") +@dataclasses.dataclass(frozen=True) +class JoinDomainPrefilterOptions: + """ + Configuration for the logical join-domain prefilter rewrite. + + Pass ``None`` to ``StreamingExecutor(join_domain_prefilter=...)`` to + disable the rewrite. + + These options can be configured via environment variables with the prefix + ``CUDF_POLARS__EXECUTOR__JOIN_DOMAIN_PREFILTER__``. + + Parameters + ---------- + threshold + Row-count ratio (domain / target) below which a derived key-domain + semi-join filter is inserted. Default is 0.5. + trace + Whether to emit plan-time trace decisions for derived key-domain + prefilters. Default is False. + """ + + _env_prefix = "CUDF_POLARS__EXECUTOR__JOIN_DOMAIN_PREFILTER" + + threshold: float = dataclasses.field( + default_factory=_make_default_factory( + f"{_env_prefix}__THRESHOLD", float, default=0.5 + ) + ) + trace: bool = dataclasses.field( + default_factory=_make_default_factory( + f"{_env_prefix}__TRACE", _bool_converter, default=False + ) + ) + + def __post_init__(self) -> None: # noqa: D105 + threshold = self.threshold + if isinstance(threshold, bool) or not isinstance(threshold, (int, float)): + raise TypeError("threshold must be a float or int") + threshold = float(threshold) + object.__setattr__(self, "threshold", threshold) + if not 0.0 <= threshold <= 1.0: + raise ValueError("threshold must be between 0 and 1") + if not isinstance(self.trace, bool): + raise TypeError("trace must be a bool") + + @dataclasses.dataclass(frozen=True, eq=True) class MemoryResourceConfig: """ @@ -640,6 +687,10 @@ class StreamingExecutor: dynamic_planning Options controlling dynamic shuffle planning. See :class:`~cudf_polars.utils.config.DynamicPlanningOptions` for more. + join_domain_prefilter + Options controlling the logical join-domain prefilter rewrite. See + :class:`~cudf_polars.utils.config.JoinDomainPrefilterOptions` for more. + ``None`` disables the rewrite. max_io_threads Maximum number of IO threads. Default is 4. This controls the parallelism of IO operations when reading data. @@ -703,6 +754,9 @@ class StreamingExecutor: dynamic_planning: DynamicPlanningOptions | None = dataclasses.field( default_factory=DynamicPlanningOptions ) + join_domain_prefilter: JoinDomainPrefilterOptions | None = dataclasses.field( + default_factory=JoinDomainPrefilterOptions + ) max_io_threads: int = dataclasses.field( default_factory=_make_default_factory( f"{_env_prefix}__MAX_IO_THREADS", int, default=4 @@ -758,6 +812,20 @@ def __post_init__(self) -> None: # noqa: D105 DynamicPlanningOptions(**self.dynamic_planning), ) + if isinstance(self.join_domain_prefilter, dict): + object.__setattr__( + self, + "join_domain_prefilter", + JoinDomainPrefilterOptions(**self.join_domain_prefilter), + ) + if self.join_domain_prefilter is not None and not isinstance( + self.join_domain_prefilter, JoinDomainPrefilterOptions + ): + raise TypeError( + "join_domain_prefilter must be a JoinDomainPrefilterOptions " + "instance, dict, or None" + ) + if self.cluster in ("spmd", "ray", "dask"): if self.sink_to_directory is False: raise ValueError( @@ -790,6 +858,7 @@ def __hash__(self) -> int: # noqa: D105 # to json and hash that. d = dataclasses.asdict(self) d["dynamic_planning"] = json.dumps(d["dynamic_planning"]) + d["join_domain_prefilter"] = json.dumps(d["join_domain_prefilter"]) return hash(tuple(sorted(d.items()))) @@ -909,6 +978,17 @@ def from_polars_engine( if not _bool_converter(env_dynamic_planning): user_executor_options["dynamic_planning"] = None + # Handle join_domain_prefilter: check user config, then env var + user_join_domain_prefilter = user_executor_options.get( + "join_domain_prefilter", None + ) + if user_join_domain_prefilter is None: + env_join_domain_prefilter = os.environ.get( + "CUDF_POLARS__EXECUTOR__JOIN_DOMAIN_PREFILTER", "1" + ) + if not _bool_converter(env_join_domain_prefilter): + user_executor_options["join_domain_prefilter"] = None + executor = StreamingExecutor(**user_executor_options) case _: # pragma: no cover; Unreachable raise ValueError(f"Unsupported executor: {user_executor}") diff --git a/python/cudf_polars/tests/streaming/test_join_domain_prefilter.py b/python/cudf_polars/tests/streaming/test_join_domain_prefilter.py new file mode 100644 index 000000000000..b978938239ae --- /dev/null +++ b/python/cudf_polars/tests/streaming/test_join_domain_prefilter.py @@ -0,0 +1,676 @@ +# 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, Literal + +import pytest + +import polars as pl + +from cudf_polars import Translator +from cudf_polars.containers import DataType +from cudf_polars.dsl import expr +from cudf_polars.dsl.ir import Cache, Filter, GroupBy, Join, Scan, Select +from cudf_polars.dsl.traversal import traversal +from cudf_polars.engine.default_singleton_engine import DefaultSingletonEngine +from cudf_polars.streaming.base import StatsCollector +from cudf_polars.streaming.join_domain_prefilter import ( + _smallest_node_containing_all, + optimize_join_domain_prefilters, +) +from cudf_polars.streaming.statistics import collect_statistics +from cudf_polars.testing.asserts import assert_gpu_result_equal +from cudf_polars.utils.config import ConfigOptions, ParquetOptions + +if TYPE_CHECKING: + import concurrent.futures + + from cudf_polars.dsl.ir import IR + from cudf_polars.streaming.base import SerializedDataSourceInfo + +I64 = DataType(pl.Int64()) +BOOL = DataType(pl.Boolean()) + + +class _SourceInfo: + type: Literal["parquet"] = "parquet" + + def __init__(self, row_count: int | None) -> None: + self.row_count = row_count + + def column_storage_size(self, column: str) -> int | None: + del column + return None + + def serialize(self) -> SerializedDataSourceInfo: + return {"type": self.type, "row_count": self.row_count, "per_file_means": {}} + + @classmethod + def deserialize(cls, data: SerializedDataSourceInfo) -> _SourceInfo: + return cls(data["row_count"]) + + +def _scan(name: str, columns: tuple[str, ...], *, predicate: bool = False) -> Scan: + schema = dict.fromkeys(columns, I64) + mask = ( + expr.NamedExpr("__predicate", expr.Literal(BOOL, True)) # noqa: FBT003 + if predicate + else None + ) + return Scan( + schema, + "parquet", + {}, + None, + [f"/tmp/{name}.parquet"], + list(columns), + 0, + -1, + None, + None, + mask, + ParquetOptions(), + ) + + +def _key(node: IR, name: str) -> expr.NamedExpr: + return expr.NamedExpr(name, expr.Col(node.schema[name], name)) + + +def _select(node: IR, **columns: str) -> Select: + schema = {output: node.schema[source] for output, source in columns.items()} + return Select( + schema, + tuple( + expr.NamedExpr(output, expr.Col(schema[output], source)) + for output, source in columns.items() + ), + True, # noqa: FBT003 + node, + ) + + +def _join( + left: IR, + right: IR, + left_on: tuple[str, ...], + right_on: tuple[str, ...], + *, + how: str = "Inner", + maintain_order: str = "none", +) -> Join: + schema = dict(left.schema) + schema.update(right.schema) + return Join( + schema, + tuple(_key(left, name) for name in left_on), + tuple(_key(right, name) for name in right_on), + (how, False, None, "_right", False, maintain_order), + left, + right, + ) + + +def _stats(**row_counts: tuple[Scan, int]) -> StatsCollector: + stats = StatsCollector() + for scan, rows in row_counts.values(): + stats.scan_stats[scan] = _SourceInfo(rows) + return stats + + +def _config( + *, dynamic_planning: bool = True, join_domain_prefilter: bool = True +) -> ConfigOptions: + executor_options: dict[str, object] = { + "join_domain_prefilter": {"trace": False} if join_domain_prefilter else None + } + if not dynamic_planning: + executor_options["dynamic_planning"] = None + return ConfigOptions.from_polars_engine( + pl.GPUEngine( + executor="streaming", + executor_options=executor_options, + ) + ) + + +def _joins(ir: IR, how: str | None = None) -> list[Join]: + return [ + node + for node in traversal([ir]) + if isinstance(node, Join) and (how is None or node.options[0] == how) + ] + + +def _contains_node(ir: IR, needle: IR) -> bool: + return any(node is needle for node in traversal([ir])) + + +def _join_key_names(keys: tuple[expr.NamedExpr, ...]) -> tuple[str, ...]: + names = [] + for key in keys: + assert isinstance(key.value, expr.Col) + names.append(key.value.name) + return tuple(names) + + +def _filtered_groupby_domain(source: IR, key: str) -> Select: + grouped = GroupBy( + {key: source.schema[key]}, + (_key(source, key),), + (), + False, # noqa: FBT003 + None, + source, + ) + filtered = Filter( + grouped.schema, + expr.NamedExpr("__predicate", expr.Literal(BOOL, True)), # noqa: FBT003 + grouped, + ) + return Select( + {key: filtered.schema[key]}, + (_key(filtered, key),), + True, # noqa: FBT003 + filtered, + ) + + +def test_simple_domain_prefilter_filters_large_side() -> None: + part = _scan("part", ("p_partkey",), predicate=True) + lineitem = _scan("lineitem", ("l_partkey", "l_suppkey")) + root = _join(part, lineitem, ("p_partkey",), ("l_partkey",)) + + optimized = optimize_join_domain_prefilters( + root, + _stats(part=(part, 6), lineitem=(lineitem, 1_800)), + _config(), + ) + + assert isinstance(optimized, Join) + assert optimized.options[0] == "Inner" + assert isinstance(optimized.children[1], Join) + assert optimized.children[1].options[0] == "Semi" + assert optimized.children[1].children[0] is lineitem + assert optimized.children[0] is part + + +def test_domain_prefilter_is_independent_of_dynamic_planning() -> None: + part = _scan("part", ("p_partkey",), predicate=True) + lineitem = _scan("lineitem", ("l_partkey",)) + root = _join(part, lineitem, ("p_partkey",), ("l_partkey",)) + + optimized = optimize_join_domain_prefilters( + root, + _stats(part=(part, 6), lineitem=(lineitem, 1_800)), + _config(dynamic_planning=False), + ) + + assert _joins(optimized, "Semi") + + +def test_domain_prefilter_can_be_disabled() -> None: + part = _scan("part", ("p_partkey",), predicate=True) + lineitem = _scan("lineitem", ("l_partkey",)) + root = _join(part, lineitem, ("p_partkey",), ("l_partkey",)) + + optimized = optimize_join_domain_prefilters( + root, + _stats(part=(part, 6), lineitem=(lineitem, 1_800)), + _config(join_domain_prefilter=False), + ) + + assert optimized is root + + +@pytest.mark.parametrize( + "nulls_equal", [False, True], ids=["nulls_not_equal", "nulls_equal"] +) +def test_nullable_join_keys_preserve_results( + nulls_equal: bool, # noqa: FBT001 + parquet_stats_executor: concurrent.futures.ThreadPoolExecutor, +) -> None: + domain = pl.LazyFrame( + { + "key": [None, 1, 2, 9], + "active": [True, True, True, False], + } + ).filter("active") + target = pl.LazyFrame( + { + "key": [None, 1, 2, 3] * 10, + "value": range(40), + } + ) + query = domain.join(target, on="key", nulls_equal=nulls_equal) + engine = pl.GPUEngine( + executor="streaming", + raise_on_fail=True, + executor_options={"join_domain_prefilter": {"threshold": 0.5}}, + ) + + ir = Translator(query._ldf.visit(), engine).translate_ir() + config = ConfigOptions.from_polars_engine(engine) + optimized = optimize_join_domain_prefilters( + ir, + collect_statistics(ir, config, parquet_stats_executor), + config, + ) + + semi_joins = _joins(optimized, "Semi") + assert semi_joins + assert all(join.options[1] is nulls_equal for join in semi_joins) + try: + assert_gpu_result_equal(query, engine=engine, check_row_order=False) + finally: + DefaultSingletonEngine.shutdown() + + +def test_no_simple_domain_prefilter_when_domain_is_not_selective() -> None: + supplier = _scan("supplier", ("s_suppkey",)) + lineitem = _scan("lineitem", ("l_suppkey",)) + root = _join(supplier, lineitem, ("s_suppkey",), ("l_suppkey",)) + + optimized = optimize_join_domain_prefilters( + root, + _stats(supplier=(supplier, 30), lineitem=(lineitem, 1_800)), + _config(), + ) + + assert optimized is root + assert not _joins(optimized, "Semi") + + +def test_composite_domain_prefilter_constrains_domain_first() -> None: + nation = _scan("nation", ("n_nationkey",), predicate=True) + orders = _scan("orders", ("o_orderkey", "n_nationkey")) + lineitem = _scan("lineitem", ("l_orderkey", "l_suppkey")) + supplier = _scan("supplier", ("s_suppkey", "s_nationkey")) + + nation_orders = _join(nation, orders, ("n_nationkey",), ("n_nationkey",)) + order_lineitem = _join( + nation_orders, + lineitem, + ("o_orderkey",), + ("l_orderkey",), + maintain_order="left", + ) + root = _join( + order_lineitem, + supplier, + ("l_suppkey", "n_nationkey"), + ("s_suppkey", "s_nationkey"), + ) + + optimized = optimize_join_domain_prefilters( + root, + _stats( + nation=(nation, 5), + orders=(orders, 900), + lineitem=(lineitem, 1_800), + supplier=(supplier, 30), + ), + _config(), + ) + + semis = _joins(optimized, "Semi") + assert isinstance(optimized, Join) + assert optimized.options[0] == "Inner" + assert optimized.children[1] is supplier + assert any(semi.children[0] is supplier for semi in semis) + assert any(semi.children[0] is lineitem for semi in semis) + + +def test_prefilter_uses_cheaper_source_domain_and_skips_expensive_domain() -> None: + part = _scan("part", ("p_partkey",), predicate=True) + partsupp = _scan("partsupp", ("ps_partkey", "ps_suppkey")) + supplier = _scan("supplier", ("s_suppkey",)) + lineitem = _scan("lineitem", ("l_partkey", "l_suppkey", "l_orderkey")) + orders = _scan("orders", ("o_orderkey",)) + + part_partsupp = _join(part, partsupp, ("p_partkey",), ("ps_partkey",)) + part_partsupp_supplier = _join( + part_partsupp, supplier, ("ps_suppkey",), ("s_suppkey",) + ) + q9_left = _join( + part_partsupp_supplier, + lineitem, + ("p_partkey", "ps_suppkey"), + ("l_partkey", "l_suppkey"), + ) + root = _join(q9_left, orders, ("l_orderkey",), ("o_orderkey",)) + + optimized = optimize_join_domain_prefilters( + root, + _stats( + part=(part, 60), + partsupp=(partsupp, 120), + supplier=(supplier, 30), + lineitem=(lineitem, 1_800), + orders=(orders, 900), + ), + _config(), + ) + + semis = _joins(optimized, "Semi") + lineitem_semis = [semi for semi in semis if semi.children[0] is lineitem] + assert lineitem_semis + assert not any(semi.children[0] is orders for semi in semis) + assert _contains_node(lineitem_semis[0].children[1], part) + assert not _contains_node(lineitem_semis[0].children[1], supplier) + + +def test_source_only_domain_does_not_stack_on_prefiltered_source() -> None: + part = _scan("part", ("p_partkey",), predicate=True) + lineitem = _scan("lineitem", ("l_partkey", "l_orderkey")) + orders = _scan("orders", ("o_orderkey",), predicate=True) + + part_lineitem = _join(part, lineitem, ("p_partkey",), ("l_partkey",)) + root = _join(part_lineitem, orders, ("l_orderkey",), ("o_orderkey",)) + + optimized = optimize_join_domain_prefilters( + root, + _stats(part=(part, 60), lineitem=(lineitem, 1_800), orders=(orders, 150)), + _config(), + ) + + lineitem_semis = [ + semi for semi in _joins(optimized, "Semi") if semi.children[0] is lineitem + ] + assert any( + _join_key_names(semi.left_on) == ("l_partkey",) for semi in lineitem_semis + ) + assert not any( + _join_key_names(semi.left_on) == ("l_orderkey",) for semi in lineitem_semis + ) + + +def test_derived_selectivity_propagates_through_rewritten_children() -> None: + region = _scan("region", ("r_regionkey",), predicate=True) + nation = _scan("nation", ("n_nationkey", "n_regionkey")) + customer = _scan("customer", ("c_custkey", "c_nationkey")) + orders = _scan("orders", ("o_orderkey", "o_custkey")) + + region_nation = _join(region, nation, ("r_regionkey",), ("n_regionkey",)) + nation_customer = _join(region_nation, customer, ("n_nationkey",), ("c_nationkey",)) + root = _join(nation_customer, orders, ("c_custkey",), ("o_custkey",)) + + optimized = optimize_join_domain_prefilters( + root, + _stats( + region=(region, 1), + nation=(nation, 25), + customer=(customer, 150), + orders=(orders, 1_500), + ), + _config(), + ) + + filtered = {semi.children[0] for semi in _joins(optimized, "Semi")} + assert {nation, customer, orders} <= filtered + + +def test_rewritten_analysis_respects_stack_and_cost_guards() -> None: + part = _scan("part", ("p_partkey",), predicate=True) + lineitem = _scan("lineitem", ("l_orderkey", "l_partkey", "l_suppkey")) + supplier = _scan("supplier", ("s_suppkey",)) + orders = _scan("orders", ("o_orderkey",), predicate=True) + + part_lineitem = _join(part, lineitem, ("p_partkey",), ("l_partkey",)) + line_supplier = _join(part_lineitem, supplier, ("l_suppkey",), ("s_suppkey",)) + root = _join(line_supplier, orders, ("l_orderkey",), ("o_orderkey",)) + + optimized = optimize_join_domain_prefilters( + root, + _stats( + part=(part, 60), + lineitem=(lineitem, 1_800), + supplier=(supplier, 30), + orders=(orders, 150), + ), + _config(), + ) + + semis = _joins(optimized, "Semi") + assert sum(semi.children[0] is lineitem for semi in semis) == 1 + assert not any(semi.children[0] is orders for semi in semis) + assert not any( + isinstance(semi.children[0], Join) and semi.children[0].options[0] == "Semi" + for semi in semis + ) + + +def test_target_source_follows_join_key_through_rename() -> None: + big = _scan("big", ("left_key", "other")) + renamed_big = _select(big, foo="left_key", other="other") + small = _scan("small", ("left_key", "other2")) + joined = _join( + renamed_big, + small, + ("other",), + ("other2",), + maintain_order="left", + ) + domain = _scan("domain", ("domain_key",), predicate=True) + root = _join(joined, domain, ("left_key",), ("domain_key",)) + + optimized = optimize_join_domain_prefilters( + root, + _stats(big=(big, 1_000), small=(small, 100), domain=(domain, 5)), + _config(), + ) + + semis = _joins(optimized, "Semi") + assert any(semi.children[0] is small for semi in semis) + assert not any(semi.children[0] is big for semi in semis) + + +def test_domain_source_follows_join_key_through_rename() -> None: + target = _scan("target", ("target_key",)) + unrelated = _scan("unrelated", ("domain_key", "other"), predicate=True) + renamed_unrelated = _select(unrelated, foo="domain_key", other="other") + domain_source = _scan("domain_source", ("domain_key", "other2"), predicate=True) + domain = _join( + renamed_unrelated, + domain_source, + ("other",), + ("other2",), + maintain_order="left", + ) + root = _join(target, domain, ("target_key",), ("domain_key",)) + + optimized = optimize_join_domain_prefilters( + root, + _stats( + target=(target, 1_000), + unrelated=(unrelated, 1), + domain_source=(domain_source, 5), + ), + _config(), + ) + + semi = next( + semi for semi in _joins(optimized, "Semi") if semi.children[0] is target + ) + selected_domain = semi.children[1] + assert isinstance(selected_domain, Select) + assert selected_domain.children[0] is domain_source + + +def test_composite_domain_columns_follow_renames() -> None: + source = _scan("source", ("raw_key", "raw_constraint")) + renamed = _select( + source, + domain_key="raw_key", + domain_constraint="raw_constraint", + ) + + producer = _smallest_node_containing_all( + renamed, + ("domain_key", "domain_constraint"), + {renamed: 20, source: 10}, + {renamed: 10, source: 10}, + ) + + assert producer is not None + assert producer.node is source + assert producer.columns == ("raw_key", "raw_constraint") + + +def test_target_replacement_does_not_rewrite_shared_domain_side() -> None: + shared = _scan("shared", ("target_key", "other")) + domain_source = _scan("domain_source", ("domain_key", "other2"), predicate=True) + domain = _join( + shared, + domain_source, + ("other",), + ("other2",), + maintain_order="left", + ) + root = _join(shared, domain, ("target_key",), ("domain_key",)) + + optimized = optimize_join_domain_prefilters( + root, + _stats(shared=(shared, 1_000), domain_source=(domain_source, 5)), + _config(), + ) + + assert isinstance(optimized, Join) + assert isinstance(optimized.children[0], Join) + assert optimized.children[0].options[0] == "Semi" + assert optimized.children[0].children[0] is shared + assert optimized.children[1] is domain + assert domain.children[0] is shared + + +def test_reuses_existing_semi_domain_to_filter_shared_target() -> None: + lineitem = _scan("lineitem", ("l_orderkey", "l_quantity")) + cached_lineitem = Cache(lineitem.schema, 0, 2, lineitem) + selected_orders = _filtered_groupby_domain(cached_lineitem, "l_orderkey") + orders = _scan("orders", ("o_orderkey", "o_custkey")) + + filtered_orders = _join( + orders, selected_orders, ("o_orderkey",), ("l_orderkey",), how="Semi" + ) + root = _join(filtered_orders, cached_lineitem, ("o_orderkey",), ("l_orderkey",)) + + optimized = optimize_join_domain_prefilters( + root, + _stats(lineitem=(lineitem, 1_800), orders=(orders, 450)), + _config(), + ) + + assert isinstance(optimized, Join) + assert optimized.options[0] == "Inner" + assert optimized.children[0] is filtered_orders + assert isinstance(optimized.children[1], Join) + assert optimized.children[1].options[0] == "Semi" + assert optimized.children[1].children[0] is cached_lineitem + assert _contains_node(optimized.children[1].children[1], selected_orders) + + +def test_reused_semi_domain_follows_filtered_key_rename() -> None: + lineitem = _scan("lineitem", ("l_orderkey", "l_quantity")) + cached_lineitem = Cache(lineitem.schema, 0, 2, lineitem) + selected_orders = _filtered_groupby_domain(cached_lineitem, "l_orderkey") + orders = _scan("orders", ("o_orderkey", "o_custkey")) + + filtered_orders = _join( + orders, selected_orders, ("o_orderkey",), ("l_orderkey",), how="Semi" + ) + renamed_orders = _select( + filtered_orders, + joined_orderkey="o_orderkey", + o_custkey="o_custkey", + ) + root = _join( + renamed_orders, + cached_lineitem, + ("joined_orderkey",), + ("l_orderkey",), + ) + + optimized = optimize_join_domain_prefilters( + root, + _stats(lineitem=(lineitem, 1_800), orders=(orders, 450)), + _config(), + ) + + assert isinstance(optimized, Join) + assert isinstance(optimized.children[1], Join) + assert optimized.children[1].options[0] == "Semi" + assert optimized.children[1].children[0] is cached_lineitem + assert _contains_node(optimized.children[1].children[1], selected_orders) + + +def test_reused_semi_domain_does_not_duplicate_uncached_source() -> None: + lineitem = _scan("lineitem", ("l_orderkey", "l_quantity")) + selected_orders = _filtered_groupby_domain(lineitem, "l_orderkey") + orders = _scan("orders", ("o_orderkey", "o_custkey")) + + filtered_orders = _join( + orders, selected_orders, ("o_orderkey",), ("l_orderkey",), how="Semi" + ) + root = _join(filtered_orders, lineitem, ("o_orderkey",), ("l_orderkey",)) + + optimized = optimize_join_domain_prefilters( + root, + _stats(lineitem=(lineitem, 1_800), orders=(orders, 450)), + _config(), + ) + + lineitem_semis = [ + semi for semi in _joins(optimized, "Semi") if semi.children[0] is lineitem + ] + assert not lineitem_semis + + +def test_reused_semi_domain_does_not_duplicate_wrapped_uncached_source() -> None: + lineitem = _scan("lineitem", ("l_orderkey", "l_quantity")) + projected_lineitem = _select( + lineitem, + l_orderkey="l_orderkey", + l_quantity="l_quantity", + ) + selected_orders = _filtered_groupby_domain(projected_lineitem, "l_orderkey") + orders = _scan("orders", ("o_orderkey", "o_custkey")) + + filtered_orders = _join( + orders, selected_orders, ("o_orderkey",), ("l_orderkey",), how="Semi" + ) + root = _join( + filtered_orders, + projected_lineitem, + ("o_orderkey",), + ("l_orderkey",), + ) + + optimized = optimize_join_domain_prefilters( + root, + _stats(lineitem=(lineitem, 1_800), orders=(orders, 450)), + _config(), + ) + + lineitem_semis = [ + semi + for semi in _joins(optimized, "Semi") + if semi.children[0] is projected_lineitem + ] + assert not lineitem_semis + + +def test_no_domain_prefilter_for_outer_join() -> None: + part = _scan("part", ("p_partkey",), predicate=True) + lineitem = _scan("lineitem", ("l_partkey",)) + root = _join(part, lineitem, ("p_partkey",), ("l_partkey",), how="Left") + + optimized = optimize_join_domain_prefilters( + root, + _stats(part=(part, 6), lineitem=(lineitem, 1_800)), + _config(), + ) + + assert optimized is root + assert not _joins(optimized, "Semi") diff --git a/python/cudf_polars/tests/streaming/test_options.py b/python/cudf_polars/tests/streaming/test_options.py index c5a42062c9a7..c9d07f687ea4 100644 --- a/python/cudf_polars/tests/streaming/test_options.py +++ b/python/cudf_polars/tests/streaming/test_options.py @@ -83,6 +83,11 @@ def test_executor_options_sink_to_directory_absent_when_unspecified() -> None: assert "sink_to_directory" not in StreamingOptions().to_executor_options() +def test_executor_options_join_domain_prefilter_disabled() -> None: + result = StreamingOptions(join_domain_prefilter=None).to_executor_options() + assert result["join_domain_prefilter"] is None + + # --------------------------------------------------------------------------- # to_engine_options # --------------------------------------------------------------------------- diff --git a/python/cudf_polars/tests/test_config.py b/python/cudf_polars/tests/test_config.py index ada8c727dd34..fcfd7771f9dc 100644 --- a/python/cudf_polars/tests/test_config.py +++ b/python/cudf_polars/tests/test_config.py @@ -29,6 +29,7 @@ Cluster, ConfigOptions, DynamicPlanningOptions, + JoinDomainPrefilterOptions, MemoryResourceConfig, StreamingExecutor, ) @@ -556,6 +557,9 @@ def test_dynamic_planning_defaults() -> None: assert config.executor.dynamic_planning.join_prefilter_threshold == 0.5 assert config.executor.dynamic_planning.join_prefilter_max_key_columns == 1 assert not config.executor.dynamic_planning.join_prefilter_trace + assert config.executor.join_domain_prefilter is not None + assert config.executor.join_domain_prefilter.threshold == 0.5 + assert not config.executor.join_domain_prefilter.trace def test_dynamic_planning_disabled_from_env(monkeypatch: pytest.MonkeyPatch) -> None: @@ -595,6 +599,31 @@ def test_join_prefilter_options_from_env(monkeypatch: pytest.MonkeyPatch) -> Non assert config.executor.dynamic_planning.join_prefilter_threshold == 0.25 assert config.executor.dynamic_planning.join_prefilter_max_key_columns is None assert config.executor.dynamic_planning.join_prefilter_trace + assert config.executor.join_domain_prefilter is not None + assert config.executor.join_domain_prefilter.threshold == 0.5 + assert not config.executor.join_domain_prefilter.trace + + +def test_join_domain_prefilter_options_from_env( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv( + "CUDF_POLARS__EXECUTOR__JOIN_DOMAIN_PREFILTER__THRESHOLD", "0.125" + ) + monkeypatch.setenv("CUDF_POLARS__EXECUTOR__JOIN_DOMAIN_PREFILTER__TRACE", "1") + config = ConfigOptions.from_polars_engine(pl.GPUEngine()) + assert config.executor.join_domain_prefilter is not None + assert config.executor.join_domain_prefilter.threshold == 0.125 + assert config.executor.join_domain_prefilter.trace + + +def test_join_domain_prefilter_disabled_from_env( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("CUDF_POLARS__EXECUTOR__JOIN_DOMAIN_PREFILTER", "0") + monkeypatch.setenv("CUDF_POLARS__EXECUTOR__JOIN_DOMAIN_PREFILTER__TRACE", "1") + config = ConfigOptions.from_polars_engine(pl.GPUEngine()) + assert config.executor.join_domain_prefilter is None @pytest.mark.parametrize("value, expected", [("none", None), ("null", None), ("2", 2)]) @@ -689,6 +718,65 @@ def test_validate_join_prefilter_trace() -> None: ) +def test_validate_join_domain_prefilter_options() -> None: + with pytest.raises(TypeError, match="threshold must be"): + ConfigOptions.from_polars_engine( + pl.GPUEngine( + executor="streaming", + executor_options={"join_domain_prefilter": {"threshold": "bad"}}, + ) + ) + with pytest.raises(ValueError, match="threshold must be between"): + ConfigOptions.from_polars_engine( + pl.GPUEngine( + executor="streaming", + executor_options={"join_domain_prefilter": {"threshold": 1.5}}, + ) + ) + with pytest.raises(TypeError, match="trace must be"): + ConfigOptions.from_polars_engine( + pl.GPUEngine( + executor="streaming", + executor_options={"join_domain_prefilter": {"trace": "bad"}}, + ) + ) + + +def test_validate_join_domain_prefilter_type() -> None: + with pytest.raises( + TypeError, + match="join_domain_prefilter must be a JoinDomainPrefilterOptions instance", + ): + ConfigOptions.from_polars_engine( + pl.GPUEngine( + executor="streaming", + executor_options={"join_domain_prefilter": object()}, + ) + ) + + +def test_join_domain_prefilter_from_instance() -> None: + options = JoinDomainPrefilterOptions(threshold=0.25, trace=True) + config = ConfigOptions.from_polars_engine( + pl.GPUEngine( + executor="streaming", + executor_options={"join_domain_prefilter": options}, + ) + ) + assert config.executor.join_domain_prefilter is options + + +def test_join_domain_prefilter_disabled_from_options() -> None: + config = ConfigOptions.from_polars_engine( + pl.GPUEngine( + executor="streaming", + executor_options={"join_domain_prefilter": None}, + ) + ) + assert config.executor.join_domain_prefilter is None + assert hash(config) == hash(config) + + def test_dynamic_planning_from_instance() -> None: config = ConfigOptions.from_polars_engine( pl.GPUEngine(