diff --git a/core.py b/core.py index 25d8232..d5b65f6 100644 --- a/core.py +++ b/core.py @@ -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, @@ -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( diff --git a/transforms/scalar/algebraic_simplify.py b/transforms/scalar/algebraic_simplify.py index 76c94af..82554a1 100644 --- a/transforms/scalar/algebraic_simplify.py +++ b/transforms/scalar/algebraic_simplify.py @@ -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"] diff --git a/transforms/scalar/constant_fold.py b/transforms/scalar/constant_fold.py index 187cff4..f628684 100644 --- a/transforms/scalar/constant_fold.py +++ b/transforms/scalar/constant_fold.py @@ -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.