Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 35 additions & 6 deletions core.py
Original file line number Diff line number Diff line change
Expand Up @@ -1516,11 +1516,38 @@ class PatternRewritePass(BasePass):
for the actual pattern matching. Iterates until convergence (no more matches).
"""

def __init__(self, pattern, rewriter, name=None, optimizer_alias=None):
# Use iterative mode - run until convergence
def __init__(self, pattern=None, rewriter=None, name=None, optimizer_alias=None, patterns=None):
"""
Initialize a pattern-rewrite pass.

Args:
pattern: Single pattern to match (backward compatibility)
rewriter: Rewriter function for the single pattern (backward compatibility)
name: Pass name
optimizer_alias: Alias for node naming
patterns: List of patterns, or list of (pattern, rewriter) tuples.
If a pattern is provided without a rewriter in the tuple,
the default 'rewriter' argument is used.
"""
super().__init__(name, optimizer_alias, iterative=True, max_iterations=100)
self.pattern = pattern
self.rewriter = trace_transformation(rewriter)

self.patterns = []

# Handle single pattern + rewriter (backward compatibility)
if pattern is not None:
self.patterns.append((pattern, trace_transformation(rewriter)))

# Handle list of patterns
if patterns is not None:
for p in patterns:
if isinstance(p, tuple):
# (pattern, rewriter) tuple
self.patterns.append((p[0], trace_transformation(p[1])))
else:
# Single pattern, use default rewriter
if rewriter is None:
raise ValueError(f"No rewriter provided for pattern: {p}")
self.patterns.append((p, trace_transformation(rewriter)))

def transform_once(
self,
Expand All @@ -1534,9 +1561,11 @@ def transform_once(
Returns:
int: Number of changes made
"""
# Register the pattern (clear first to avoid duplicates)
# Register all patterns (clear first to avoid duplicates)
optimizer.clear_transformations()
optimizer.add_transformation(self.pattern, self.rewriter)

for p, r in self.patterns:
optimizer.add_transformation(p, r)

# Run one pattern matching iteration
new_graph_def, changes = optimizer.match_patterns_once(
Expand Down
13 changes: 10 additions & 3 deletions transforms/scalar/algebraic_simplify.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,9 +87,16 @@ class AlgebraicSimplifyPass(PatternRewritePass):
"""

def __init__(self):
# We'll handle multiple patterns manually in _rewrite
pattern = Any(alias="op") # fallback, we check inside
super().__init__(pattern, self._rewrite, name="AlgebraicSimplify")
# We register specific Op patterns instead of Any() to enable O(1) matching.
# This significantly improves performance on large graphs by avoiding
# the wildcard matching path for every node.
supported_ops = [
"Add", "Sub", "Mul", "Div", "Neg", "LogicalNot", "Abs", "Square",
"Sqrt", "Pow", "Equal", "NotEqual", "Less", "Greater", "LessEqual",
"GreaterEqual", "LogicalAnd", "LogicalOr", "Select", "Identity"
]
patterns = [Op(op, alias="op") for op in supported_ops]
super().__init__(patterns=patterns, rewriter=self._rewrite, name="AlgebraicSimplify")

def _rewrite(self, match, optimizer):
node = match.matched_nodes["op"]
Expand Down
15 changes: 12 additions & 3 deletions transforms/scalar/constant_fold.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,9 +58,18 @@ class ConstantFoldPass(PatternRewritePass):
"""

def __init__(self):
# Matches any operation with all inputs as Const
pattern = Any(alias="op")
super().__init__(pattern, self._rewrite_constant_op, name="ConstantFold")
# We register specific Op patterns instead of Any() to enable O(1) matching.
supported_ops = [
"Add", "Mul", "Sub", "Div", "Neg", "Equal", "NotEqual", "Less",
"Greater", "LessEqual", "GreaterEqual", "LogicalAnd", "LogicalOr",
"LogicalNot", "BitwiseAnd", "BitwiseOr", "BitwiseXor", "Abs",
"Exp", "Expm1", "Log", "Log1p", "Sqrt", "Pow", "Rsqrt", "Square",
"Sin", "Cos", "Tan", "Asin", "Acos", "Atan", "Atan2", "Floor",
"Ceil", "Round", "Sign", "Reshape", "Transpose", "ConcatV2",
"Select", "Cast"
]
patterns = [Op(op, alias="op") for op in supported_ops]
super().__init__(patterns=patterns, rewriter=self._rewrite_constant_op, name="ConstantFold")

def _is_all_const(self, inputs, optimizer):
"""Check if all inputs are Const nodes.
Expand Down