diff --git a/bigframes/core/__init__.py b/bigframes/core/__init__.py index 04291edbb1..185ce7cd4f 100644 --- a/bigframes/core/__init__.py +++ b/bigframes/core/__init__.py @@ -183,7 +183,7 @@ def project_to_id(self, expression: ex.Expression, output_id: str): child=self.node, assignments=tuple(exprs), ) - ) + ).merge_projections() def assign(self, source_id: str, destination_id: str) -> ArrayValue: if destination_id in self.column_ids: # Mutate case @@ -208,7 +208,7 @@ def assign(self, source_id: str, destination_id: str) -> ArrayValue: child=self.node, assignments=tuple(exprs), ) - ) + ).merge_projections() def assign_constant( self, @@ -242,7 +242,7 @@ def assign_constant( child=self.node, assignments=tuple(exprs), ) - ) + ).merge_projections() def select_columns(self, column_ids: typing.Sequence[str]) -> ArrayValue: selections = ((ex.free_var(col_id), col_id) for col_id in column_ids) @@ -251,7 +251,7 @@ def select_columns(self, column_ids: typing.Sequence[str]) -> ArrayValue: child=self.node, assignments=tuple(selections), ) - ) + ).merge_projections() def drop_columns(self, columns: Iterable[str]) -> ArrayValue: new_projection = ( @@ -264,7 +264,7 @@ def drop_columns(self, columns: Iterable[str]) -> ArrayValue: child=self.node, assignments=tuple(new_projection), ) - ) + ).merge_projections() def aggregate( self, @@ -466,3 +466,7 @@ def _uniform_sampling(self, fraction: float) -> ArrayValue: The row numbers of result is non-deterministic, avoid to use. """ return ArrayValue(nodes.RandomSampleNode(self.node, fraction)) + + def merge_projections(self) -> ArrayValue: + new_node = bigframes.core.rewrite.maybe_squash_projection(self.node) + return ArrayValue(new_node) diff --git a/bigframes/core/compile/compiled.py b/bigframes/core/compile/compiled.py index 88c1006c79..d14a5d3241 100644 --- a/bigframes/core/compile/compiled.py +++ b/bigframes/core/compile/compiled.py @@ -1050,8 +1050,8 @@ def _hide_column(self, column_id) -> OrderedIR: def _bake_ordering(self) -> OrderedIR: """Bakes ordering expression into the selection, maybe creating hidden columns.""" ordering_expressions = self._ordering.all_ordering_columns - new_exprs = [] - new_baked_cols = [] + new_exprs: list[OrderingExpression] = [] + new_baked_cols: list[ibis_types.Value] = [] for expr in ordering_expressions: if isinstance(expr.scalar_expression, ex.OpExpression): baked_column = self._compile_expression(expr.scalar_expression).name( @@ -1059,18 +1059,28 @@ def _bake_ordering(self) -> OrderedIR: ) new_baked_cols.append(baked_column) new_expr = OrderingExpression( - ex.free_var(baked_column.name), expr.direction, expr.na_last + ex.free_var(baked_column.get_name()), expr.direction, expr.na_last ) new_exprs.append(new_expr) - else: + elif isinstance(expr.scalar_expression, ex.UnboundVariableExpression): + order_col = expr.scalar_expression.id new_exprs.append(expr) - - ordering = self._ordering.with_ordering_columns(new_exprs) + if order_col not in self.column_ids: + new_baked_cols.append( + self._ibis_bindings[expr.scalar_expression.id] + ) + + new_ordering = ExpressionOrdering( + tuple(new_exprs), + self._ordering.integer_encoding, + self._ordering.string_encoding, + self._ordering.total_ordering_columns, + ) return OrderedIR( self._table, columns=self.columns, - hidden_ordering_columns=[*self._hidden_ordering_columns, *new_baked_cols], - ordering=ordering, + hidden_ordering_columns=tuple(new_baked_cols), + ordering=new_ordering, predicates=self._predicates, ) diff --git a/bigframes/core/rewrite.py b/bigframes/core/rewrite.py index 61fe28b7b5..e3a07c04b4 100644 --- a/bigframes/core/rewrite.py +++ b/bigframes/core/rewrite.py @@ -35,16 +35,21 @@ class SquashedSelect: columns: Tuple[Tuple[scalar_exprs.Expression, str], ...] predicate: Optional[scalar_exprs.Expression] ordering: Tuple[order.OrderingExpression, ...] + reverse_root: bool = False @classmethod - def from_node(cls, node: nodes.BigFrameNode) -> SquashedSelect: + def from_node( + cls, node: nodes.BigFrameNode, projections_only: bool = False + ) -> SquashedSelect: if isinstance(node, nodes.ProjectionNode): - return cls.from_node(node.child).project(node.assignments) - elif isinstance(node, nodes.FilterNode): + return cls.from_node(node.child, projections_only=projections_only).project( + node.assignments + ) + elif not projections_only and isinstance(node, nodes.FilterNode): return cls.from_node(node.child).filter(node.predicate) - elif isinstance(node, nodes.ReversedNode): + elif not projections_only and isinstance(node, nodes.ReversedNode): return cls.from_node(node.child).reverse() - elif isinstance(node, nodes.OrderByNode): + elif not projections_only and isinstance(node, nodes.OrderByNode): return cls.from_node(node.child).order_with(node.by) else: selection = tuple( @@ -63,7 +68,9 @@ def project( new_columns = tuple( (expr.bind_all_variables(self.column_lookup), id) for expr, id in projection ) - return SquashedSelect(self.root, new_columns, self.predicate, self.ordering) + return SquashedSelect( + self.root, new_columns, self.predicate, self.ordering, self.reverse_root + ) def filter(self, predicate: scalar_exprs.Expression) -> SquashedSelect: if self.predicate is None: @@ -72,18 +79,24 @@ def filter(self, predicate: scalar_exprs.Expression) -> SquashedSelect: new_predicate = ops.and_op.as_expr( self.predicate, predicate.bind_all_variables(self.column_lookup) ) - return SquashedSelect(self.root, self.columns, new_predicate, self.ordering) + return SquashedSelect( + self.root, self.columns, new_predicate, self.ordering, self.reverse_root + ) def reverse(self) -> SquashedSelect: new_ordering = tuple(expr.with_reverse() for expr in self.ordering) - return SquashedSelect(self.root, self.columns, self.predicate, new_ordering) + return SquashedSelect( + self.root, self.columns, self.predicate, new_ordering, not self.reverse_root + ) def order_with(self, by: Tuple[order.OrderingExpression, ...]): adjusted_orderings = [ order_part.bind_variables(self.column_lookup) for order_part in by ] new_ordering = (*adjusted_orderings, *self.ordering) - return SquashedSelect(self.root, self.columns, self.predicate, new_ordering) + return SquashedSelect( + self.root, self.columns, self.predicate, new_ordering, self.reverse_root + ) def maybe_join( self, right: SquashedSelect, join_def: join_defs.JoinDefinition @@ -126,8 +139,10 @@ def maybe_join( new_columns = remap_names(join_def, lselection, rselection) # Reconstruct ordering + reverse_root = self.reverse_root if join_type == "right": new_ordering = right.ordering + reverse_root = right.reverse_root elif join_type == "outer": if lmask is not None: prefix = order.OrderingExpression(lmask, order.OrderingDirection.DESC) @@ -158,11 +173,15 @@ def maybe_join( new_ordering = self.ordering else: raise ValueError(f"Unexpected join type {join_type}") - return SquashedSelect(self.root, new_columns, new_predicate, new_ordering) + return SquashedSelect( + self.root, new_columns, new_predicate, new_ordering, reverse_root + ) def expand(self) -> nodes.BigFrameNode: # Safest to apply predicates first, as it may filter out inputs that cannot be handled by other expressions root = self.root + if self.reverse_root: + root = nodes.ReversedNode(child=root) if self.predicate: root = nodes.FilterNode(child=root, predicate=self.predicate) if self.ordering: @@ -170,6 +189,15 @@ def expand(self) -> nodes.BigFrameNode: return nodes.ProjectionNode(child=root, assignments=self.columns) +def maybe_squash_projection(node: nodes.BigFrameNode) -> nodes.BigFrameNode: + if isinstance(node, nodes.ProjectionNode) and isinstance( + node.child, nodes.ProjectionNode + ): + # Conservative approach, only squash consecutive projections, even though could also squash filters, reorderings + return SquashedSelect.from_node(node, projections_only=True).expand() + return node + + def maybe_rewrite_join(join_node: nodes.JoinNode) -> nodes.BigFrameNode: left_side = SquashedSelect.from_node(join_node.left_child) right_side = SquashedSelect.from_node(join_node.right_child)