Source code for ray.data._internal.logical.operators.n_ary_operator

import enum
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, List, Optional

from ray.data._internal.logical.interfaces import (
    LogicalOperator,
    LogicalOperatorSupportsPredicatePassThrough,
    LogicalOperatorUnifiesInputSchemas,
    PredicatePassThroughBehavior,
)
from ray.util.annotations import PublicAPI

if TYPE_CHECKING:
    from ray.data.block import Schema

__all__ = [
    "Mix",
    "MixStoppingCondition",
    "NAry",
    "Union",
    "Zip",
]


[docs] @PublicAPI(stability="alpha") class MixStoppingCondition(enum.Enum): """Controls when a mix pipeline terminates. STOP_ON_SHORTEST: Pipeline ends when the shortest dataset is exhausted. Other datasets are truncated. STOP_ON_LONGEST_DROP: Pipeline ends when the longest dataset is exhausted. Shorter datasets drop out once exhausted; later batches are drawn entirely from longer datasets. """ STOP_ON_SHORTEST = "stop_on_shortest" STOP_ON_LONGEST_DROP = "stop_on_longest_drop"
def estimate_num_mix_outputs( per_input_counts: List[Optional[int]], weights: List[float], stopping_condition: MixStoppingCondition, ) -> Optional[int]: """Estimate total output count for a mix operation. Used by both the logical and physical Mix operators to estimate num_outputs_total / num_output_rows_total. """ if any(c is None for c in per_input_counts): return None if stopping_condition == MixStoppingCondition.STOP_ON_LONGEST_DROP: return sum(per_input_counts) elif stopping_condition == MixStoppingCondition.STOP_ON_SHORTEST: # Limited by whichever input runs out first relative to its weight. total_weight = sum(weights) return min( int(count / (w / total_weight)) for count, w in zip(per_input_counts, weights) ) else: raise ValueError(f"Unknown stopping condition: {stopping_condition}") @dataclass(frozen=True, repr=False, eq=False, init=False) class NAry(LogicalOperator): """Base class for n-ary operators, which take multiple input operators.""" def __init__( self, input_dependencies: List[LogicalOperator], ): """Initialize the n-ary operator. Args: input_dependencies: The input operators. """ object.__setattr__(self, "_input_dependencies", list(input_dependencies)) def _with_new_input_dependencies( self, input_dependencies: List[LogicalOperator] ) -> LogicalOperator: return self.__class__(input_dependencies) @dataclass(frozen=True, repr=False, eq=False, init=False) class Zip(NAry): """Logical operator for zip.""" _input_dependencies: List[LogicalOperator] = field(init=False, repr=False) def __init__( self, input_dependencies: List[LogicalOperator], ): for input_op in input_dependencies: assert isinstance(input_op, LogicalOperator), input_op object.__setattr__(self, "_input_dependencies", list(input_dependencies)) def estimated_num_outputs(self): total_num_outputs = 0 for input in self.input_dependencies: num_outputs = input.estimated_num_outputs() if num_outputs is None: return None total_num_outputs = max(total_num_outputs, num_outputs) return total_num_outputs def infer_schema(self) -> Optional["Schema"]: # Reuse the runtime ``BlockAccessor.zip`` so plan-time and # execution-time schemas agree by construction (same column # suffixing rules, etc.). import pyarrow as pa from ray.data.block import BlockAccessor input_schemas = [op.infer_schema() for op in self.input_dependencies] if not input_schemas or not all( isinstance(s, pa.Schema) for s in input_schemas ): return None try: combined = input_schemas[0].empty_table() for s in input_schemas[1:]: combined = BlockAccessor.for_block(combined).zip(s.empty_table()) except (pa.ArrowTypeError, pa.ArrowInvalid): return None return combined.schema @dataclass(frozen=True, repr=False, eq=False, init=False) class Mix(NAry, LogicalOperatorUnifiesInputSchemas): """Logical operator for weighted dataset mixing.""" _name: str = field(init=False, repr=False) _input_dependencies: List[LogicalOperator] = field(init=False, repr=False) weights: List[float] = field(init=False, repr=False) stopping_condition: MixStoppingCondition = field(init=False, repr=False) def __init__( self, input_dependencies: List[LogicalOperator], *, weights: List[float], stopping_condition: MixStoppingCondition = MixStoppingCondition.STOP_ON_SHORTEST, ): if len(input_dependencies) != len(weights): raise ValueError( f"Number of input operators ({len(input_dependencies)}) must match " f"number of weights ({len(weights)})." ) if any(weight <= 0 for weight in weights): raise ValueError(f"Weights must be positive. Got weights: {weights}") for input_op in input_dependencies: assert isinstance(input_op, LogicalOperator), input_op object.__setattr__(self, "_name", self.__class__.__name__) object.__setattr__(self, "_input_dependencies", list(input_dependencies)) object.__setattr__(self, "weights", weights) object.__setattr__(self, "stopping_condition", stopping_condition) def estimated_num_outputs(self) -> Optional[int]: if self.stopping_condition == MixStoppingCondition.STOP_ON_SHORTEST: return None return estimate_num_mix_outputs( [op.estimated_num_outputs() for op in self.input_dependencies], self.weights, self.stopping_condition, ) def _with_new_input_dependencies( self, input_dependencies: List[LogicalOperator] ) -> LogicalOperator: return self.__class__( input_dependencies, weights=self.weights, stopping_condition=self.stopping_condition, ) @dataclass(frozen=True, repr=False, eq=False, init=False) class Union( NAry, LogicalOperatorSupportsPredicatePassThrough, LogicalOperatorUnifiesInputSchemas, ): """Logical operator for union.""" _input_dependencies: List[LogicalOperator] = field(init=False, repr=False) def __init__( self, input_dependencies: List[LogicalOperator], ): for input_op in input_dependencies: assert isinstance(input_op, LogicalOperator), input_op object.__setattr__(self, "_input_dependencies", list(input_dependencies)) def estimated_num_outputs(self): total_num_outputs = 0 for input in self.input_dependencies: num_outputs = input.estimated_num_outputs() if num_outputs is None: return None total_num_outputs += num_outputs return total_num_outputs def predicate_passthrough_behavior(self) -> PredicatePassThroughBehavior: # Union allows pushing filter into each branch return PredicatePassThroughBehavior.PUSH_INTO_BRANCHES