diff --git a/benchmarks/benchmark_wildcard.py b/benchmarks/benchmark_wildcard.py new file mode 100644 index 0000000..76e56da --- /dev/null +++ b/benchmarks/benchmark_wildcard.py @@ -0,0 +1,43 @@ +import time +import tensorflow.compat.v1 as tf +import numpy as np +from graph_optimizer.core import GraphOptimizer, Op, Any +from graph_optimizer.transforms.scalar.algebraic_simplify import AlgebraicSimplifyPass +from graph_optimizer.transforms.scalar.constant_fold import ConstantFoldPass +from graph_optimizer.utils.graph_utils import create_node, create_const_node + +def create_large_sparse_graph(num_nodes=5000): + graph_def = tf.GraphDef() + # Create many placeholders + for i in range(num_nodes): + p = create_node("Placeholder", f"p_{i}", attr={"dtype": tf.AttrValue(type=tf.float32.as_datatype_enum)}) + graph_def.node.extend([p]) + + # Create some Add(x, 0) nodes scattered around + zero = create_const_node("zero", value=0, dtype="float32", shape=[]) + graph_def.node.extend([zero]) + + for i in range(100): + idx = i * (num_nodes // 100) + a = create_node("Add", f"add_{i}", inputs=[f"p_{idx}", "zero"]) + graph_def.node.extend([a]) + + return graph_def + +def benchmark(): + num_nodes = 100000 + print(f"Generating graph with {num_nodes} nodes...") + graph_def = create_large_sparse_graph(num_nodes) + + optimizer = GraphOptimizer(graph_def) + pass_obj = AlgebraicSimplifyPass() + + print("Running AlgebraicSimplifyPass benchmark...") + start_time = time.time() + optimizer.optimize(pass_name="AlgebraicSimplify", max_iterations=1) + end_time = time.time() + + print(f"Time taken: {end_time - start_time:.4f}s") + +if __name__ == "__main__": + benchmark() diff --git a/core.py b/core.py index 25d8232..c9d92f1 100644 --- a/core.py +++ b/core.py @@ -1,6 +1,7 @@ import tensorflow.compat.v1 as tf import collections import time +import itertools from typing import Dict, List, Set, Optional, Any as AnyType, Tuple, Union from .utils.logger import ( logger as logging, @@ -582,7 +583,7 @@ def match_once( if node.name in replaced_node_names: continue - candidates = self.pattern_index.get(node.op, []) + self.wildcard_patterns + candidates = itertools.chain(self.pattern_index.get(node.op, []), self.wildcard_patterns) found_match = False for pattern, rewriter in candidates: @@ -1516,11 +1517,46 @@ 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, + patterns: Union[Pattern, List[Union[Pattern, Tuple[Pattern, AnyType]]]] = None, + rewriter: AnyType = None, + name=None, + optimizer_alias=None, + pattern: Pattern = None, # Backward compatibility for keyword argument + ): + """ + Initialize a pattern-rewrite pass. + + Args: + patterns: List of (Pattern, Rewriter) tuples, or a single Pattern (for backward compatibility). + rewriter: Single Rewriter (only used if patterns is a single Pattern). + name: Pass name. + optimizer_alias: Optimizer alias. + pattern: Single Pattern (backward compatibility for keyword argument). + """ super().__init__(name, optimizer_alias, iterative=True, max_iterations=100) - self.pattern = pattern - self.rewriter = trace_transformation(rewriter) + + # Handle backward compatibility for 'pattern' keyword argument + if pattern is not None: + patterns = pattern + + self.patterns = [] + # Support both new 'patterns' list and old 'pattern, rewriter' arguments + if isinstance(patterns, list): + # Cache wrapped rewriters to avoid redundant wrappers for multiple patterns + rewriter_to_wrapped = {} + for p in patterns: + if isinstance(p, tuple): + pat, rew = p + if rew not in rewriter_to_wrapped: + rewriter_to_wrapped[rew] = trace_transformation(rew) + self.patterns.append((pat, rewriter_to_wrapped[rew])) + else: + self.patterns.append(p) + elif patterns is not None and rewriter is not None: + # Single pattern and rewriter case + self.patterns = [(patterns, trace_transformation(rewriter))] def transform_once( self, @@ -1534,9 +1570,10 @@ 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..a7ca035 100644 --- a/transforms/scalar/algebraic_simplify.py +++ b/transforms/scalar/algebraic_simplify.py @@ -87,9 +87,102 @@ 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") + # Register specific Op patterns to leverage indexed matching + 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"), self._rewrite) for op in supported_ops] + super().__init__(patterns=patterns, name="AlgebraicSimplify") + + def _mapped_result(self, name, target_name): + """Returns a RewriteResult that redirects a node to another.""" + return RewriteResult(new_nodes=[], node_mapping={name: target_name}) + + def _new_node_result(self, name, new_node): + """Returns a RewriteResult that replaces a node with a new one.""" + return RewriteResult(new_nodes=[new_node], node_mapping={name: new_node.name}) + + def _bool_const(self, name, val, inputs, optimizer): + """Helper for comparison results. Result has same shape as first input.""" + s = self._get_shape(inputs[0], optimizer) + if s is None: + return None + return self._new_node_result( + name, create_const_node(name + "_bool", value=val, dtype="bool", shape=s) + ) + + def _get_node(self, name, optimizer): + """Helper to get node object ignoring output index.""" + real_name = name.split(":")[0] + return optimizer.nodes.get(real_name) + + def _is_const(self, node_name, value, optimizer): + """Helper to check if a node is Const with given value (broadcast-safe).""" + node = self._get_node(node_name, optimizer) + if node is None or node.op != "Const": + return False + val = optimizer.get_node_attr(node, "value") + # Check if all elements are equal to the target value + return np.all(np.equal(val, value)) + + def _get_shape(self, node_name, optimizer): + """Helper to get shape of a node.""" + node = self._get_node(node_name, optimizer) + if node is None: + return None + # Check for shape attribute (Placeholder, etc.) + if "shape" in node.attr: + return [d.size for d in node.attr["shape"].shape.dim] + # Check for Const value shape + if node.op == "Const" and "value" in node.attr: + tensor = node.attr["value"].tensor + if tensor.HasField("tensor_shape"): + return [d.size for d in tensor.tensor_shape.dim] + return None + + def _is_scalar(self, node_name, optimizer): + """Helper to check if a node is definitely scalar.""" + shape = self._get_shape(node_name, optimizer) + return shape == [] + + def _get_broadcast_shape(self, s1, s2): + """Helper to compute broadcast shape of two shapes.""" + if s1 is None or s2 is None: + return None + if s1 == s2: + return s1 + if not s1: + return s2 + if not s2: + return s1 + + # Simple broadcasting logic + len1, len2 = len(s1), len(s2) + max_len = max(len1, len2) + result = [] + for i in range(max_len): + d1 = s1[len1 - 1 - i] if i < len1 else 1 + d2 = s2[len2 - 1 - i] if i < len2 else 1 + if d1 == d2: + result.append(d1) + elif d1 == 1: + result.append(d2) + elif d2 == 1: + result.append(d1) + else: + return None # Incompatible + return result[::-1] + + def _is_shape_preserving(self, source_shape, target_shape): + """Helper to check if simplification is shape-preserving.""" + # If both are unknown, assume it's safe (common in simple tests) + if source_shape is None and target_shape is None: + return True + if source_shape is None or target_shape is None: + return False + return source_shape == target_shape def _rewrite(self, match, optimizer): node = match.matched_nodes["op"] @@ -97,303 +190,200 @@ def _rewrite(self, match, optimizer): inputs = list(node.input) name = node.name - def _mapped_result(target_name): - return RewriteResult(new_nodes=[], node_mapping={name: target_name}) - - def _new_node_result(new_node): - return RewriteResult( - new_nodes=[new_node], node_mapping={name: new_node.name} - ) - - # Helper to create True/False const - def _bool_const(val): - return _new_node_result( - create_const_node(name + "_bool", value=val, dtype="bool", shape=[]) - ) - - # Helper to get node object ignoring output index - def _get_node(name): - real_name = name.split(":")[0] - return optimizer.nodes.get(real_name) - - # Helper to check if a node is Const with given value (broadcast-safe) - def _is_const(node_name, value): - node = _get_node(node_name) - if node is None: - return False - if node.op != "Const": - return False - val = optimizer.get_node_attr(node, "value") - # Check if all elements are equal to the target value - return np.all(np.equal(val, value)) - - # Helper to get shape of a node - def _get_shape(node_name): - node = _get_node(node_name) - if node is None: - return None - # Check for shape attribute (Placeholder, etc.) - if "shape" in node.attr: - return [d.size for d in node.attr["shape"].shape.dim] - # Check for Const value shape - if node.op == "Const" and "value" in node.attr: - tensor = node.attr["value"].tensor - if tensor.HasField("tensor_shape"): - return [d.size for d in tensor.tensor_shape.dim] - return None - - # Helper to check if a node is definitely scalar - def _is_scalar(node_name): - shape = _get_shape(node_name) - return shape == [] - - # Helper to compute broadcast shape of two shapes - def _get_broadcast_shape(s1, s2): - if s1 is None or s2 is None: - return None - if s1 == s2: - return s1 - if not s1: - return s2 - if not s2: - return s1 - - # Simple broadcasting logic - len1, len2 = len(s1), len(s2) - max_len = max(len1, len2) - result = [] - for i in range(max_len): - d1 = s1[len1 - 1 - i] if i < len1 else 1 - d2 = s2[len2 - 1 - i] if i < len2 else 1 - if d1 == d2: - result.append(d1) - elif d1 == 1: - result.append(d2) - elif d2 == 1: - result.append(d1) - else: - return None # Incompatible - return result[::-1] - - # Helper to check if simplification is shape-preserving - def _is_shape_preserving(source_shape, target_shape): - # If both are unknown, assume it's safe (common in simple tests) - if source_shape is None and target_shape is None: - return True - if source_shape is None or target_shape is None: - return False - return source_shape == target_shape - # Rule: Add(x, 0) or Add(0, x) if op_type == "Add": left, right = inputs[0], inputs[1] - s_left, s_right = _get_shape(left), _get_shape(right) - s_res = _get_broadcast_shape(s_left, s_right) + s_left, s_right = self._get_shape(left, optimizer), self._get_shape(right, optimizer) + s_res = self._get_broadcast_shape(s_left, s_right) - if _is_const(left, 0) and _is_shape_preserving(s_res, s_right): - return _mapped_result(right) - if _is_const(right, 0) and _is_shape_preserving(s_res, s_left): - return _mapped_result(left) + if self._is_const(left, 0, optimizer) and self._is_shape_preserving(s_res, s_right): + return self._mapped_result(name, right) + if self._is_const(right, 0, optimizer) and self._is_shape_preserving(s_res, s_left): + return self._mapped_result(name, left) # Add(x, Neg(x)) -> 0 or Add(Neg(x), x) -> 0 - # Note: This is a simplified check for Neg(x) for l, r in [(left, right), (right, left)]: - rn = _get_node(r) + rn = self._get_node(r, optimizer) if rn and rn.op == "Neg" and rn.input[0] == l: - s = _get_shape(l) + s = self._get_shape(l, optimizer) if s is not None: - source = _get_node(l) + source = self._get_node(l, optimizer) dtype = source.attr.get("dtype", "float32") if source else "float32" - return _new_node_result( - create_const_node(name + "_zero", value=0, dtype=dtype, shape=s) + return self._new_node_result( + name, create_const_node(name + "_zero", value=0, dtype=dtype, shape=s) ) # Rule: Sub(x, 0) → x if op_type == "Sub": left, right = inputs[0], inputs[1] - if _is_const(right, 0) and ( - _is_scalar(right) or _get_shape(right) == _get_shape(left) + if self._is_const(right, 0, optimizer) and ( + self._is_scalar(right, optimizer) or self._get_shape(right, optimizer) == self._get_shape(left, optimizer) ): - return _mapped_result(left) + return self._mapped_result(name, left) # Sub(x, x) → 0 if left == right: - s = _get_shape(left) + s = self._get_shape(left, optimizer) if s is not None: - source = _get_node(left) + source = self._get_node(left, optimizer) dtype = source.attr.get("dtype", "float32") if source else "float32" - return _new_node_result( - create_const_node(name + "_zero", value=0, dtype=dtype, shape=s) + return self._new_node_result( + name, create_const_node(name + "_zero", value=0, dtype=dtype, shape=s) ) # Rule: Mul(x, 1) or Mul(1, x) if op_type == "Mul": left, right = inputs[0], inputs[1] - s_left, s_right = _get_shape(left), _get_shape(right) - s_res = _get_broadcast_shape(s_left, s_right) + s_left, s_right = self._get_shape(left, optimizer), self._get_shape(right, optimizer) + s_res = self._get_broadcast_shape(s_left, s_right) - if _is_const(left, 1) and _is_shape_preserving(s_res, s_right): - return _mapped_result(right) - if _is_const(right, 1) and _is_shape_preserving(s_res, s_left): - return _mapped_result(left) + if self._is_const(left, 1, optimizer) and self._is_shape_preserving(s_res, s_right): + return self._mapped_result(name, right) + if self._is_const(right, 1, optimizer) and self._is_shape_preserving(s_res, s_left): + return self._mapped_result(name, left) # Mul(x, 0) → 0 - if _is_const(left, 0) or _is_const(right, 0): + if self._is_const(left, 0, optimizer) or self._is_const(right, 0, optimizer): if s_res is not None: - source_name = right if _is_const(left, 0) else left - source = _get_node(source_name) + source_name = right if self._is_const(left, 0, optimizer) else left + source = self._get_node(source_name, optimizer) dtype = source.attr.get("dtype", "float32") if source else "float32" - return _new_node_result( - create_const_node( + return self._new_node_result( + name, create_const_node( name + "_zero", value=0, dtype=dtype, shape=s_res ) ) # Mul(x, x) -> Square(x) if left == right: - return _new_node_result( - create_node("Square", name + "_sq", inputs=[left]) + return self._new_node_result( + name, create_node("Square", name + "_sq", inputs=[left]) ) # Rule: Div(x, 1) → x if op_type == "Div": left, right = inputs[0], inputs[1] - s_left, s_right = _get_shape(left), _get_shape(right) - s_res = _get_broadcast_shape(s_left, s_right) - if _is_const(right, 1) and _is_shape_preserving(s_res, s_left): - return _mapped_result(left) + s_left, s_right = self._get_shape(left, optimizer), self._get_shape(right, optimizer) + s_res = self._get_broadcast_shape(s_left, s_right) + if self._is_const(right, 1, optimizer) and self._is_shape_preserving(s_res, s_left): + return self._mapped_result(name, left) # Div(x, x) -> 1 if left == right: - s = _get_shape(left) + s = self._get_shape(left, optimizer) if s is not None: - source = _get_node(left) + source = self._get_node(left, optimizer) dtype = source.attr.get("dtype", "float32") if source else "float32" - return _new_node_result( - create_const_node(name + "_one", value=1, dtype=dtype, shape=s) + return self._new_node_result( + name, create_const_node(name + "_one", value=1, dtype=dtype, shape=s) ) # Rule: Neg(Neg(x)) → x if op_type == "Neg": - inp = _get_node(inputs[0]) + inp = self._get_node(inputs[0], optimizer) if inp and inp.op == "Neg": - return _mapped_result(inp.input[0]) + return self._mapped_result(name, inp.input[0]) # Rule: LogicalNot(LogicalNot(x)) → x if op_type == "LogicalNot": - inp = _get_node(inputs[0]) + inp = self._get_node(inputs[0], optimizer) if inp and inp.op == "LogicalNot": - return _mapped_result(inp.input[0]) + return self._mapped_result(name, inp.input[0]) # Rule: Abs(Abs(x)) → Abs(x) if op_type == "Abs": - inp = _get_node(inputs[0]) + inp = self._get_node(inputs[0], optimizer) if inp and inp.op == "Abs": - orig = _get_node(inp.input[0]) + orig = self._get_node(inp.input[0], optimizer) if orig: - return _new_node_result( - create_node("Abs", name + "_abs", inputs=[orig.name]) + return self._new_node_result( + name, create_node("Abs", name + "_abs", inputs=[orig.name]) ) # Rule: Square(Sqrt(x)) → x (domain assumed ok) if op_type == "Square": - inp = _get_node(inputs[0]) + inp = self._get_node(inputs[0], optimizer) if inp and inp.op == "Sqrt": - return _mapped_result(inp.input[0]) + return self._mapped_result(name, inp.input[0]) # Rule: Sqrt(Square(x)) → Abs(x) if op_type == "Sqrt": - inp = _get_node(inputs[0]) + inp = self._get_node(inputs[0], optimizer) if inp and inp.op == "Square": - orig = _get_node(inp.input[0]) + orig = self._get_node(inp.input[0], optimizer) if orig: - return _new_node_result( - create_node("Abs", name + "_abs", inputs=[orig.name]) + return self._new_node_result( + name, create_node("Abs", name + "_abs", inputs=[orig.name]) ) # Rule: Pow(x, 1) -> x if op_type == "Pow": left, right = inputs[0], inputs[1] - s_left, s_right = _get_shape(left), _get_shape(right) - s_res = _get_broadcast_shape(s_left, s_right) - if _is_const(right, 1) and _is_shape_preserving(s_res, s_left): - return _mapped_result(left) + s_left, s_right = self._get_shape(left, optimizer), self._get_shape(right, optimizer) + s_res = self._get_broadcast_shape(s_left, s_right) + if self._is_const(right, 1, optimizer) and self._is_shape_preserving(s_res, s_left): + return self._mapped_result(name, left) # Pow(x, 2) -> Square(x) - if _is_const(right, 2) and _is_shape_preserving(s_res, s_left): - return _new_node_result( - create_node("Square", name + "_sq", inputs=[left]) + if self._is_const(right, 2, optimizer) and self._is_shape_preserving(s_res, s_left): + return self._new_node_result( + name, create_node("Square", name + "_sq", inputs=[left]) ) - # Helper for comparison results - def _comparison_const(val): - # Equal(x, x) -> True should have same shape as x (or broadcasted shape) - # If x is [2, 2], result is [2, 2] of True - s = _get_shape(inputs[0]) - if s is None: - return None # Safer to skip if shape unknown - return _new_node_result( - create_const_node(name + "_bool", value=val, dtype="bool", shape=s) - ) - # Rule: Equal(x, x) → True - if op_type == "Equal": - left, right = inputs[0], inputs[1] - if left == right: - return _comparison_const(True) + if op_type == "Equal" and inputs[0] == inputs[1]: + return self._bool_const(name, True, inputs, optimizer) # Rule: NotEqual(x, x) → False - if op_type == "NotEqual": - left, right = inputs[0], inputs[1] - if left == right: - return _comparison_const(False) + if op_type == "NotEqual" and inputs[0] == inputs[1]: + return self._bool_const(name, False, inputs, optimizer) # Rule: Less(x, x) → False ; Greater(x, x) → False if op_type in ("Less", "Greater") and inputs[0] == inputs[1]: - return _comparison_const(False) + return self._bool_const(name, False, inputs, optimizer) # Rule: LessEqual(x, x) → True ; GreaterEqual(x, x) → True if op_type in ("LessEqual", "GreaterEqual") and inputs[0] == inputs[1]: - return _comparison_const(True) + return self._bool_const(name, True, inputs, optimizer) # Rule: And(x, True) → x ; And(True, x) → x if op_type == "LogicalAnd": left, right = inputs[0], inputs[1] - s_left, s_right = _get_shape(left), _get_shape(right) - s_res = _get_broadcast_shape(s_left, s_right) + s_left, s_right = self._get_shape(left, optimizer), self._get_shape(right, optimizer) + s_res = self._get_broadcast_shape(s_left, s_right) - if _is_const(left, True) and _is_shape_preserving(s_res, s_right): - return _mapped_result(right) - if _is_const(right, True) and _is_shape_preserving(s_res, s_left): - return _mapped_result(left) + if self._is_const(left, True, optimizer) and self._is_shape_preserving(s_res, s_right): + return self._mapped_result(name, right) + if self._is_const(right, True, optimizer) and self._is_shape_preserving(s_res, s_left): + return self._mapped_result(name, left) # LogicalAnd(x, x) -> x if left == right: - return _mapped_result(left) + return self._mapped_result(name, left) # LogicalAnd(x, False) -> False - if _is_const(left, False) or _is_const(right, False): + if self._is_const(left, False, optimizer) or self._is_const(right, False, optimizer): if s_res is not None: - return _new_node_result( - create_const_node(name + "_bool", value=False, dtype="bool", shape=s_res) + return self._new_node_result( + name, create_const_node(name + "_bool", value=False, dtype="bool", shape=s_res) ) # Rule: Or(x, False) → x ; Or(False, x) → x if op_type == "LogicalOr": left, right = inputs[0], inputs[1] - s_left, s_right = _get_shape(left), _get_shape(right) - s_res = _get_broadcast_shape(s_left, s_right) + s_left, s_right = self._get_shape(left, optimizer), self._get_shape(right, optimizer) + s_res = self._get_broadcast_shape(s_left, s_right) - if _is_const(left, False) and _is_shape_preserving(s_res, s_right): - return _mapped_result(right) - if _is_const(right, False) and _is_shape_preserving(s_res, s_left): - return _mapped_result(left) + if self._is_const(left, False, optimizer) and self._is_shape_preserving(s_res, s_right): + return self._mapped_result(name, right) + if self._is_const(right, False, optimizer) and self._is_shape_preserving(s_res, s_left): + return self._mapped_result(name, left) # LogicalOr(x, x) -> x if left == right: - return _mapped_result(left) + return self._mapped_result(name, left) # LogicalOr(x, True) -> True - if _is_const(left, True) or _is_const(right, True): + if self._is_const(left, True, optimizer) or self._is_const(right, True, optimizer): if s_res is not None: - return _new_node_result( - create_const_node(name + "_bool", value=True, dtype="bool", shape=s_res) + return self._new_node_result( + name, create_const_node(name + "_bool", value=True, dtype="bool", shape=s_res) ) # Rule: Select(cond, x, x) → x if op_type == "Select": if len(inputs) >= 3 and inputs[1] == inputs[2]: - return _mapped_result(inputs[1]) + return self._mapped_result(name, inputs[1]) # Rule: Identity(x) -> x (bypass or collapse nested Identity) if op_type == "Identity": @@ -410,14 +400,14 @@ def _comparison_const(val): if "_class" in node.attr: return None # Collapse nested Identity - inp_node = _get_node(inputs[0]) + inp_node = self._get_node(inputs[0], optimizer) if inp_node and inp_node.op == "Identity": inner_input = inp_node.input[0] new_node = create_node( "Identity", name + "_collapsed", inputs=[inner_input] ) - return _new_node_result(new_node) + return self._new_node_result(name, new_node) # Bypass single Identity - return _mapped_result(inputs[0]) + return self._mapped_result(name, inputs[0]) return None diff --git a/transforms/scalar/constant_fold.py b/transforms/scalar/constant_fold.py index 187cff4..22a3af3 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") + # Define supported ops + 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" + ] + # Register specific Op patterns to leverage indexed matching + patterns = [(Op(op, alias="op"), self._rewrite_constant_op) for op in supported_ops] + super().__init__(patterns=patterns, name="ConstantFold") def _is_all_const(self, inputs, optimizer): """Check if all inputs are Const nodes. @@ -123,177 +132,45 @@ def _rewrite_constant_op(self, match, optimizer): op_type = op_node.op - # Define supported ops - def _add(x, y): - return np.add(x, y) - - def _mul(x, y): - return np.multiply(x, y) - - def _sub(x, y): - return np.subtract(x, y) - - def _div(x, y): - # Safety: check for division by zero to avoid inf/nan nodes - # If we want to allow them, we could just let numpy handle it. - # But usually it's safer to avoid folding into Inf/NaN if it might crash later. - # However, for robustness, maybe we should allow it if numpy does. - with np.errstate(divide="ignore", invalid="ignore"): - res = np.divide(x, y) - return res - - def _neg(x): - return np.negative(x) - - def _equal(x, y): - return np.equal(x, y) - - def _not_equal(x, y): - return np.not_equal(x, y) - - def _less(x, y): - return np.less(x, y) - - def _greater(x, y): - return np.greater(x, y) - - def _less_equal(x, y): - return np.less_equal(x, y) - - def _greater_equal(x, y): - return np.greater_equal(x, y) - - def _logical_and(x, y): - return np.logical_and(x, y) - - def _logical_or(x, y): - return np.logical_or(x, y) - - def _logical_not(x): - return np.logical_not(x) - - def _bitwise_and(x, y): - return np.bitwise_and(x.astype(np.int64), y.astype(np.int64)) - - def _bitwise_or(x, y): - return np.bitwise_or(x.astype(np.int64), y.astype(np.int64)) - - def _bitwise_xor(x, y): - return np.bitwise_xor(x.astype(np.int64), y.astype(np.int64)) - - def _abs(x): - return np.abs(x) - - def _exp(x): - return np.exp(x) - - def _expm1(x): - return np.expm1(x) - - def _log(x): - with np.errstate(divide="ignore", invalid="ignore"): - return np.log(x) - - def _log1p(x): - return np.log1p(x) - - def _sqrt(x): - with np.errstate(invalid="ignore"): - return np.sqrt(x) - - def _pow(x, y): - return np.power(x, y) - - def _rsqrt(x): - with np.errstate(divide="ignore", invalid="ignore"): - return 1.0 / np.sqrt(x) - - def _square(x): - return np.square(x) - - def _sin(x): - return np.sin(x) - - def _cos(x): - return np.cos(x) - - def _tan(x): - return np.tan(x) - - def _asin(x): - return np.arcsin(x) - - def _acos(x): - return np.arccos(x) - - def _atan(x): - return np.arctan(x) - - def _atan2(y, x): - return np.arctan2(y, x) - - def _floor(x): - return np.floor(x) - - def _ceil(x): - return np.ceil(x) - - def _round(x): - return np.round(x) - - def _sign(x): - return np.sign(x) - - def _reshape(x, shape): - return np.reshape(x, shape) - - def _transpose(x, axes): - return np.transpose(x, axes) - - def _concatenate(x_list, axis=0): - return np.concatenate(x_list, axis=axis) - - def _select(cond, x, y): - return np.where(cond, x, y) - + # Define operations map (using static definitions to avoid re-creation) ops_map = { - "Add": lambda: _add(*arrays[:2]), - "Mul": lambda: _mul(*arrays[:2]), - "Sub": lambda: _sub(*arrays[:2]), - "Div": lambda: _div(*arrays[:2]), - "Neg": lambda: _neg(arrays[0]), - "Equal": lambda: _equal(*arrays[:2]), - "NotEqual": lambda: _not_equal(*arrays[:2]), - "Less": lambda: _less(*arrays[:2]), - "Greater": lambda: _greater(*arrays[:2]), - "LessEqual": lambda: _less_equal(*arrays[:2]), - "GreaterEqual": lambda: _greater_equal(*arrays[:2]), - "LogicalAnd": lambda: _logical_and(*arrays[:2]), - "LogicalOr": lambda: _logical_or(*arrays[:2]), - "LogicalNot": lambda: _logical_not(arrays[0]), - "BitwiseAnd": lambda: _bitwise_and(*arrays[:2]), - "BitwiseOr": lambda: _bitwise_or(*arrays[:2]), - "BitwiseXor": lambda: _bitwise_xor(*arrays[:2]), - "Abs": lambda: _abs(arrays[0]), - "Exp": lambda: _exp(arrays[0]), - "Expm1": lambda: _expm1(arrays[0]), - "Log": lambda: _log(arrays[0]), - "Log1p": lambda: _log1p(arrays[0]), - "Sqrt": lambda: _sqrt(arrays[0]), - "Pow": lambda: _pow(*arrays[:2]), - "Rsqrt": lambda: _rsqrt(arrays[0]), - "Square": lambda: _square(arrays[0]), - "Sin": lambda: _sin(arrays[0]), - "Cos": lambda: _cos(arrays[0]), - "Tan": lambda: _tan(arrays[0]), - "Asin": lambda: _asin(arrays[0]), - "Acos": lambda: _acos(arrays[0]), - "Atan": lambda: _atan(arrays[0]), - "Atan2": lambda: _atan2(*arrays[:2]), - "Floor": lambda: _floor(arrays[0]), - "Ceil": lambda: _ceil(arrays[0]), - "Round": lambda: _round(arrays[0]), - "Sign": lambda: _sign(arrays[0]), + "Add": lambda: np.add(arrays[0], arrays[1]), + "Mul": lambda: np.multiply(arrays[0], arrays[1]), + "Sub": lambda: np.subtract(arrays[0], arrays[1]), + "Div": lambda: np.divide(arrays[0], arrays[1]), + "Neg": lambda: np.negative(arrays[0]), + "Equal": lambda: np.equal(arrays[0], arrays[1]), + "NotEqual": lambda: np.not_equal(arrays[0], arrays[1]), + "Less": lambda: np.less(arrays[0], arrays[1]), + "Greater": lambda: np.greater(arrays[0], arrays[1]), + "LessEqual": lambda: np.less_equal(arrays[0], arrays[1]), + "GreaterEqual": lambda: np.greater_equal(arrays[0], arrays[1]), + "LogicalAnd": lambda: np.logical_and(arrays[0], arrays[1]), + "LogicalOr": lambda: np.logical_or(arrays[0], arrays[1]), + "LogicalNot": lambda: np.logical_not(arrays[0]), + "BitwiseAnd": lambda: np.bitwise_and(arrays[0].astype(np.int64), arrays[1].astype(np.int64)), + "BitwiseOr": lambda: np.bitwise_or(arrays[0].astype(np.int64), arrays[1].astype(np.int64)), + "BitwiseXor": lambda: np.bitwise_xor(arrays[0].astype(np.int64), arrays[1].astype(np.int64)), + "Abs": lambda: np.abs(arrays[0]), + "Exp": lambda: np.exp(arrays[0]), + "Expm1": lambda: np.expm1(arrays[0]), + "Log": lambda: np.log(arrays[0]), + "Log1p": lambda: np.log1p(arrays[0]), + "Sqrt": lambda: np.sqrt(arrays[0]), + "Pow": lambda: np.power(arrays[0], arrays[1]), + "Rsqrt": lambda: 1.0 / np.sqrt(arrays[0]), + "Square": lambda: np.square(arrays[0]), + "Sin": lambda: np.sin(arrays[0]), + "Cos": lambda: np.cos(arrays[0]), + "Tan": lambda: np.tan(arrays[0]), + "Asin": lambda: np.arcsin(arrays[0]), + "Acos": lambda: np.arccos(arrays[0]), + "Atan": lambda: np.arctan(arrays[0]), + "Atan2": lambda: np.arctan2(arrays[0], arrays[1]), + "Floor": lambda: np.floor(arrays[0]), + "Ceil": lambda: np.ceil(arrays[0]), + "Round": lambda: np.round(arrays[0]), + "Sign": lambda: np.sign(arrays[0]), } # Handle special cases requiring extra attrs @@ -301,17 +178,17 @@ def _select(cond, x, y): shape_arr = arrays[1] if shape_arr.ndim != 1: return None - result = _reshape(arrays[0], tuple(shape_arr.astype(int))) + result = np.reshape(arrays[0], tuple(shape_arr.astype(int))) elif op_type == "Transpose": axes_arr = arrays[1] if axes_arr.ndim != 1: return None - result = _transpose(arrays[0], tuple(axes_arr.astype(int))) + result = np.transpose(arrays[0], tuple(axes_arr.astype(int))) elif op_type == "ConcatV2": axis_val = int(arrays[-1]) - result = _concatenate(arrays[:-1], axis=axis_val) + result = np.concatenate(arrays[:-1], axis=axis_val) elif op_type == "Select": - result = _select(*arrays[:3]) + result = np.where(arrays[0], arrays[1], arrays[2]) elif op_type == "Cast": # Cast to target dtype dst_t_attr = op_node.attr.get("DstT", None)