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
43 changes: 43 additions & 0 deletions benchmarks/benchmark_wildcard.py
Original file line number Diff line number Diff line change
@@ -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()
51 changes: 44 additions & 7 deletions core.py
Original file line number Diff line number Diff line change
@@ -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,
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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,
Expand All @@ -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(
Expand Down
Loading