metanion 0.1.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- metanion-0.1.0/PKG-INFO +57 -0
- metanion-0.1.0/README.md +35 -0
- metanion-0.1.0/metanion/__init__.py +39 -0
- metanion-0.1.0/metanion/algebra/__init__.py +17 -0
- metanion-0.1.0/metanion/algebra/rewrite_rules.py +497 -0
- metanion-0.1.0/metanion/api/__init__.py +7 -0
- metanion-0.1.0/metanion/api/metanion.py +231 -0
- metanion-0.1.0/metanion/calculus/__init__.py +16 -0
- metanion-0.1.0/metanion/calculus/derivative_rules.py +583 -0
- metanion-0.1.0/metanion/calculus/symbolic_differentiator.py +472 -0
- metanion-0.1.0/metanion/compile/__init__.py +19 -0
- metanion-0.1.0/metanion/compile/bytecode_compiler.py +95 -0
- metanion-0.1.0/metanion/compile/lazy_graph.py +355 -0
- metanion-0.1.0/metanion/compile/straight_line_program.py +17 -0
- metanion-0.1.0/metanion/config.py +117 -0
- metanion-0.1.0/metanion/core/__init__.py +12 -0
- metanion-0.1.0/metanion/core/dtype_system.py +89 -0
- metanion-0.1.0/metanion/core/memory_arena.py +169 -0
- metanion-0.1.0/metanion/core/tensor.py +191 -0
- metanion-0.1.0/metanion/core/tensor_buffer.py +247 -0
- metanion-0.1.0/metanion/core/tensor_shape.py +389 -0
- metanion-0.1.0/metanion/data/__init__.py +1 -0
- metanion-0.1.0/metanion/data/dataset.py +281 -0
- metanion-0.1.0/metanion/data/statistics_injector.py +191 -0
- metanion-0.1.0/metanion/exceptions.py +94 -0
- metanion-0.1.0/metanion/gp/__init__.py +48 -0
- metanion-0.1.0/metanion/gp/bloat_control.py +302 -0
- metanion-0.1.0/metanion/gp/crossover.py +275 -0
- metanion-0.1.0/metanion/gp/fitness.py +392 -0
- metanion-0.1.0/metanion/gp/individual.py +301 -0
- metanion-0.1.0/metanion/gp/initialization.py +189 -0
- metanion-0.1.0/metanion/gp/mutation.py +261 -0
- metanion-0.1.0/metanion/gp/population.py +220 -0
- metanion-0.1.0/metanion/gp/safe_ops.py +142 -0
- metanion-0.1.0/metanion/gp/selection.py +248 -0
- metanion-0.1.0/metanion/io/__init__.py +12 -0
- metanion-0.1.0/metanion/io/binary_decoder.py +232 -0
- metanion-0.1.0/metanion/io/binary_encoder.py +231 -0
- metanion-0.1.0/metanion/io/checkpoint_manager.py +293 -0
- metanion-0.1.0/metanion/metanion_engine.py +49 -0
- metanion-0.1.0/metanion/model/__init__.py +8 -0
- metanion-0.1.0/metanion/model/metanion_layer.py +256 -0
- metanion-0.1.0/metanion/model/metanion_model.py +357 -0
- metanion-0.1.0/metanion/model/metanion_stack.py +133 -0
- metanion-0.1.0/metanion/runtime/__init__.py +24 -0
- metanion-0.1.0/metanion/runtime/gc_controller.py +203 -0
- metanion-0.1.0/metanion/runtime/jit_cache_manager.py +143 -0
- metanion-0.1.0/metanion/symbolic/__init__.py +45 -0
- metanion-0.1.0/metanion/symbolic/expression_node.py +91 -0
- metanion-0.1.0/metanion/symbolic/handle_utils.py +74 -0
- metanion-0.1.0/metanion/symbolic/hash_consing_pool.py +101 -0
- metanion-0.1.0/metanion/symbolic/op_enum.py +176 -0
- metanion-0.1.0/metanion/symbolic/op_metadata.py +539 -0
- metanion-0.1.0/metanion/utils/__init__.py +44 -0
- metanion-0.1.0/metanion/utils/cost_model.py +336 -0
- metanion-0.1.0/metanion/utils/time_profiler.py +230 -0
- metanion-0.1.0/metanion/utils/tree_printer.py +323 -0
- metanion-0.1.0/metanion.egg-info/PKG-INFO +57 -0
- metanion-0.1.0/metanion.egg-info/SOURCES.txt +72 -0
- metanion-0.1.0/metanion.egg-info/dependency_links.txt +1 -0
- metanion-0.1.0/metanion.egg-info/requires.txt +1 -0
- metanion-0.1.0/metanion.egg-info/top_level.txt +2 -0
- metanion-0.1.0/pyproject.toml +3 -0
- metanion-0.1.0/setup.cfg +4 -0
- metanion-0.1.0/setup.py +20 -0
- metanion-0.1.0/tests/__init__.py +4 -0
- metanion-0.1.0/tests/run_all_tests.py +70 -0
- metanion-0.1.0/tests/test_all_functions.py +508 -0
- metanion-0.1.0/tests/test_calculus.py +75 -0
- metanion-0.1.0/tests/test_compile.py +69 -0
- metanion-0.1.0/tests/test_core.py +72 -0
- metanion-0.1.0/tests/test_io.py +54 -0
- metanion-0.1.0/tests/test_symbolic.py +80 -0
- metanion-0.1.0/tests/test_training.py +100 -0
metanion-0.1.0/PKG-INFO
ADDED
|
@@ -0,0 +1,57 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: metanion
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: A zero-weight symbolic tensor engine
|
|
5
|
+
Home-page: https://github.com/rohitpatraoutlook-dotcom/metanion
|
|
6
|
+
Author: Rohit Patra
|
|
7
|
+
Author-email: rohitpatra@outlook.com
|
|
8
|
+
Classifier: Programming Language :: Python :: 3
|
|
9
|
+
Classifier: License :: OSI Approved :: MIT License
|
|
10
|
+
Requires-Python: >=3.8
|
|
11
|
+
Description-Content-Type: text/markdown
|
|
12
|
+
Requires-Dist: numpy>=1.19.0
|
|
13
|
+
Dynamic: author
|
|
14
|
+
Dynamic: author-email
|
|
15
|
+
Dynamic: classifier
|
|
16
|
+
Dynamic: description
|
|
17
|
+
Dynamic: description-content-type
|
|
18
|
+
Dynamic: home-page
|
|
19
|
+
Dynamic: requires-dist
|
|
20
|
+
Dynamic: requires-python
|
|
21
|
+
Dynamic: summary
|
|
22
|
+
|
|
23
|
+
# 🧠 Metanion - Zero-Weight Symbolic Tensor Engine
|
|
24
|
+
|
|
25
|
+
[](https://badge.fury.io/py/metanion)
|
|
26
|
+
[](https://www.python.org/downloads/)
|
|
27
|
+
[](https://opensource.org/licenses/MIT)
|
|
28
|
+
|
|
29
|
+
**Metanion** is a revolutionary tensor engine where **weights are symbolic expressions, not numbers**. It learns mathematical relationships using Genetic Programming and JIT compilation.
|
|
30
|
+
|
|
31
|
+
## 🎯 Why Metanion?
|
|
32
|
+
|
|
33
|
+
- ✅ **No Numerical Weights** - Only operation sequences stored
|
|
34
|
+
- ✅ **Explainable** - Outputs human-readable equations
|
|
35
|
+
- ✅ **Fast** - JIT compiled to Python bytecode
|
|
36
|
+
- ✅ **Differentiable** - Full symbolic differentiation
|
|
37
|
+
- ✅ **Lightweight** - Minimal memory footprint
|
|
38
|
+
|
|
39
|
+
## 🚀 Quick Start
|
|
40
|
+
|
|
41
|
+
```python
|
|
42
|
+
from metanion import create_model, train, predict
|
|
43
|
+
import numpy as np
|
|
44
|
+
|
|
45
|
+
# Create data: y = 2*x + 1 + noise
|
|
46
|
+
X = np.random.randn(200, 1)
|
|
47
|
+
y = 2 * X[:, 0] + 1 + 0.1 * np.random.randn(200)
|
|
48
|
+
|
|
49
|
+
# Create and train model
|
|
50
|
+
model = create_model([1, 10, 1])
|
|
51
|
+
train(X, y, epochs=30)
|
|
52
|
+
|
|
53
|
+
# Make predictions
|
|
54
|
+
predictions = predict(X)
|
|
55
|
+
|
|
56
|
+
# Get the learned equation
|
|
57
|
+
print(model._best_individual.get_expression())
|
metanion-0.1.0/README.md
ADDED
|
@@ -0,0 +1,35 @@
|
|
|
1
|
+
# 🧠 Metanion - Zero-Weight Symbolic Tensor Engine
|
|
2
|
+
|
|
3
|
+
[](https://badge.fury.io/py/metanion)
|
|
4
|
+
[](https://www.python.org/downloads/)
|
|
5
|
+
[](https://opensource.org/licenses/MIT)
|
|
6
|
+
|
|
7
|
+
**Metanion** is a revolutionary tensor engine where **weights are symbolic expressions, not numbers**. It learns mathematical relationships using Genetic Programming and JIT compilation.
|
|
8
|
+
|
|
9
|
+
## 🎯 Why Metanion?
|
|
10
|
+
|
|
11
|
+
- ✅ **No Numerical Weights** - Only operation sequences stored
|
|
12
|
+
- ✅ **Explainable** - Outputs human-readable equations
|
|
13
|
+
- ✅ **Fast** - JIT compiled to Python bytecode
|
|
14
|
+
- ✅ **Differentiable** - Full symbolic differentiation
|
|
15
|
+
- ✅ **Lightweight** - Minimal memory footprint
|
|
16
|
+
|
|
17
|
+
## 🚀 Quick Start
|
|
18
|
+
|
|
19
|
+
```python
|
|
20
|
+
from metanion import create_model, train, predict
|
|
21
|
+
import numpy as np
|
|
22
|
+
|
|
23
|
+
# Create data: y = 2*x + 1 + noise
|
|
24
|
+
X = np.random.randn(200, 1)
|
|
25
|
+
y = 2 * X[:, 0] + 1 + 0.1 * np.random.randn(200)
|
|
26
|
+
|
|
27
|
+
# Create and train model
|
|
28
|
+
model = create_model([1, 10, 1])
|
|
29
|
+
train(X, y, epochs=30)
|
|
30
|
+
|
|
31
|
+
# Make predictions
|
|
32
|
+
predictions = predict(X)
|
|
33
|
+
|
|
34
|
+
# Get the learned equation
|
|
35
|
+
print(model._best_individual.get_expression())
|
|
@@ -0,0 +1,39 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Metanion - A Zero-Weight Symbolic Tensor Engine
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
__version__ = "0.1.0"
|
|
6
|
+
__author__ = "Metanion Team"
|
|
7
|
+
|
|
8
|
+
# Core
|
|
9
|
+
from .core import Tensor, DType, Shape, get_arena, reset_arena
|
|
10
|
+
|
|
11
|
+
# Symbolic
|
|
12
|
+
from .symbolic import (
|
|
13
|
+
OpID, OpCategory, intern, lookup, simplify,
|
|
14
|
+
get_pool, reset_pool, get_depth, count_nodes_in_subtree,
|
|
15
|
+
get_op_name, get_op_arity, is_binary_op, is_unary_op,
|
|
16
|
+
ExpressionNode, ExpressionNodeFactory, get_all_operation_ids
|
|
17
|
+
)
|
|
18
|
+
|
|
19
|
+
# Compile
|
|
20
|
+
from .compile import compile_handle, StraightLineProgram
|
|
21
|
+
|
|
22
|
+
# API - Simple interface
|
|
23
|
+
from .api import Metanion
|
|
24
|
+
|
|
25
|
+
__all__ = [
|
|
26
|
+
'__version__', '__author__',
|
|
27
|
+
'Tensor', 'DType', 'Shape',
|
|
28
|
+
'get_arena', 'reset_arena',
|
|
29
|
+
'OpID', 'OpCategory',
|
|
30
|
+
'intern', 'lookup', 'simplify',
|
|
31
|
+
'get_pool', 'reset_pool',
|
|
32
|
+
'get_depth', 'count_nodes_in_subtree',
|
|
33
|
+
'get_op_name', 'get_op_arity',
|
|
34
|
+
'is_binary_op', 'is_unary_op',
|
|
35
|
+
'ExpressionNode', 'ExpressionNodeFactory',
|
|
36
|
+
'get_all_operation_ids',
|
|
37
|
+
'compile_handle', 'StraightLineProgram',
|
|
38
|
+
'Metanion', # Main API
|
|
39
|
+
]
|
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
"""Algebraic simplification for Metanion."""
|
|
2
|
+
|
|
3
|
+
class RewriteSystem:
|
|
4
|
+
def __init__(self):
|
|
5
|
+
self._rules = []
|
|
6
|
+
def normalize(self, handle):
|
|
7
|
+
return handle
|
|
8
|
+
|
|
9
|
+
_REWRITE_SYSTEM = None
|
|
10
|
+
|
|
11
|
+
def get_rewrite_system():
|
|
12
|
+
global _REWRITE_SYSTEM
|
|
13
|
+
if _REWRITE_SYSTEM is None:
|
|
14
|
+
_REWRITE_SYSTEM = RewriteSystem()
|
|
15
|
+
return _REWRITE_SYSTEM
|
|
16
|
+
|
|
17
|
+
__all__ = ['get_rewrite_system']
|
|
@@ -0,0 +1,497 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Algebraic rewrite rules for the Metanion engine.
|
|
3
|
+
Implements term rewriting system (TRS) for expression simplification.
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
from typing import Optional, Tuple, List, Dict, Callable, Any, Set
|
|
7
|
+
from dataclasses import dataclass, field
|
|
8
|
+
from enum import Enum
|
|
9
|
+
|
|
10
|
+
from ..symbolic import OpID, get_op_arity, get_op_name, intern, lookup
|
|
11
|
+
from ..symbolic import ExpressionNode, get_pool
|
|
12
|
+
from ..exceptions import ExpressionError
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class RewriteDirection(Enum):
|
|
16
|
+
"""Direction of rewrite rule application."""
|
|
17
|
+
LEFT_TO_RIGHT = "->"
|
|
18
|
+
RIGHT_TO_LEFT = "<-"
|
|
19
|
+
BOTH = "<->"
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
@dataclass
|
|
23
|
+
class Pattern:
|
|
24
|
+
"""
|
|
25
|
+
A pattern for matching expression trees.
|
|
26
|
+
Supports variables and wildcards.
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
op: Optional[OpID] = None
|
|
30
|
+
left: Optional['Pattern'] = None
|
|
31
|
+
right: Optional['Pattern'] = None
|
|
32
|
+
is_variable: bool = False
|
|
33
|
+
var_name: Optional[str] = None
|
|
34
|
+
is_wildcard: bool = False
|
|
35
|
+
|
|
36
|
+
def __post_init__(self):
|
|
37
|
+
"""Validate the pattern."""
|
|
38
|
+
if self.is_variable and self.var_name is None:
|
|
39
|
+
self.var_name = f"_var_{id(self)}"
|
|
40
|
+
|
|
41
|
+
def match(self, handle: int, pool) -> Optional[Dict[str, int]]:
|
|
42
|
+
"""
|
|
43
|
+
Match this pattern against an expression.
|
|
44
|
+
|
|
45
|
+
Args:
|
|
46
|
+
handle: The handle to match against.
|
|
47
|
+
pool: The expression pool.
|
|
48
|
+
|
|
49
|
+
Returns:
|
|
50
|
+
A mapping of variable names to handles, or None if match fails.
|
|
51
|
+
"""
|
|
52
|
+
if self.is_wildcard:
|
|
53
|
+
return {}
|
|
54
|
+
|
|
55
|
+
if self.is_variable:
|
|
56
|
+
return {self.var_name: handle}
|
|
57
|
+
|
|
58
|
+
node = pool.get_node(handle)
|
|
59
|
+
if node is None:
|
|
60
|
+
return None
|
|
61
|
+
|
|
62
|
+
if node.op != self.op:
|
|
63
|
+
return None
|
|
64
|
+
|
|
65
|
+
# Match children
|
|
66
|
+
bindings = {}
|
|
67
|
+
children = node.get_children()
|
|
68
|
+
|
|
69
|
+
if self.left is not None:
|
|
70
|
+
if len(children) < 1:
|
|
71
|
+
return None
|
|
72
|
+
left_match = self.left.match(children[0], pool)
|
|
73
|
+
if left_match is None:
|
|
74
|
+
return None
|
|
75
|
+
bindings.update(left_match)
|
|
76
|
+
|
|
77
|
+
if self.right is not None:
|
|
78
|
+
if len(children) < 2:
|
|
79
|
+
return None
|
|
80
|
+
right_match = self.right.match(children[1], pool)
|
|
81
|
+
if right_match is None:
|
|
82
|
+
return None
|
|
83
|
+
bindings.update(right_match)
|
|
84
|
+
|
|
85
|
+
return bindings
|
|
86
|
+
|
|
87
|
+
def substitute(self, bindings: Dict[str, int], pool) -> int:
|
|
88
|
+
"""
|
|
89
|
+
Substitute variables in the pattern with values from bindings.
|
|
90
|
+
|
|
91
|
+
Args:
|
|
92
|
+
bindings: Mapping of variable names to handles.
|
|
93
|
+
pool: The expression pool.
|
|
94
|
+
|
|
95
|
+
Returns:
|
|
96
|
+
The handle of the substituted expression.
|
|
97
|
+
"""
|
|
98
|
+
if self.is_variable:
|
|
99
|
+
if self.var_name not in bindings:
|
|
100
|
+
raise ExpressionError(f"Variable {self.var_name} not bound")
|
|
101
|
+
return bindings[self.var_name]
|
|
102
|
+
|
|
103
|
+
if self.is_wildcard:
|
|
104
|
+
# Wildcard - return some default
|
|
105
|
+
return intern(OpID.CONST_ZERO)
|
|
106
|
+
|
|
107
|
+
left_handle = None
|
|
108
|
+
right_handle = None
|
|
109
|
+
|
|
110
|
+
if self.left is not None:
|
|
111
|
+
left_handle = self.left.substitute(bindings, pool)
|
|
112
|
+
|
|
113
|
+
if self.right is not None:
|
|
114
|
+
right_handle = self.right.substitute(bindings, pool)
|
|
115
|
+
|
|
116
|
+
return intern(self.op, left_handle, right_handle)
|
|
117
|
+
|
|
118
|
+
@classmethod
|
|
119
|
+
def var(cls, name: str) -> 'Pattern':
|
|
120
|
+
"""Create a variable pattern."""
|
|
121
|
+
return cls(is_variable=True, var_name=name)
|
|
122
|
+
|
|
123
|
+
@classmethod
|
|
124
|
+
def wildcard(cls) -> 'Pattern':
|
|
125
|
+
"""Create a wildcard pattern."""
|
|
126
|
+
return cls(is_wildcard=True)
|
|
127
|
+
|
|
128
|
+
@classmethod
|
|
129
|
+
def op(cls, op_id: OpID, left: Optional['Pattern'] = None, right: Optional['Pattern'] = None) -> 'Pattern':
|
|
130
|
+
"""Create an operation pattern."""
|
|
131
|
+
return cls(op=op_id, left=left, right=right)
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
@dataclass
|
|
135
|
+
class RewriteRule:
|
|
136
|
+
"""
|
|
137
|
+
A rewrite rule for transforming expressions.
|
|
138
|
+
"""
|
|
139
|
+
|
|
140
|
+
name: str
|
|
141
|
+
lhs: Pattern
|
|
142
|
+
rhs: Pattern
|
|
143
|
+
direction: RewriteDirection = RewriteDirection.LEFT_TO_RIGHT
|
|
144
|
+
priority: int = 0
|
|
145
|
+
description: str = ""
|
|
146
|
+
condition: Optional[Callable[[Dict[str, int], int], bool]] = None
|
|
147
|
+
|
|
148
|
+
def apply(self, handle: int, pool) -> Optional[int]:
|
|
149
|
+
"""
|
|
150
|
+
Apply the rewrite rule to an expression.
|
|
151
|
+
|
|
152
|
+
Args:
|
|
153
|
+
handle: The handle to rewrite.
|
|
154
|
+
pool: The expression pool.
|
|
155
|
+
|
|
156
|
+
Returns:
|
|
157
|
+
The rewritten handle, or None if the rule doesn't apply.
|
|
158
|
+
"""
|
|
159
|
+
# Match the LHS pattern
|
|
160
|
+
bindings = self.lhs.match(handle, pool)
|
|
161
|
+
if bindings is None:
|
|
162
|
+
return None
|
|
163
|
+
|
|
164
|
+
# Check the condition (if any)
|
|
165
|
+
if self.condition is not None:
|
|
166
|
+
if not self.condition(bindings, handle):
|
|
167
|
+
return None
|
|
168
|
+
|
|
169
|
+
# Substitute the RHS pattern
|
|
170
|
+
return self.rhs.substitute(bindings, pool)
|
|
171
|
+
|
|
172
|
+
def __repr__(self) -> str:
|
|
173
|
+
"""String representation."""
|
|
174
|
+
arrow = self.direction.value
|
|
175
|
+
return f"{self.name}: {self.lhs} {arrow} {self.rhs}"
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
class RewriteSystem:
|
|
179
|
+
"""
|
|
180
|
+
Term rewriting system for simplifying expressions.
|
|
181
|
+
Applies rewrite rules in order until no more rules apply.
|
|
182
|
+
"""
|
|
183
|
+
|
|
184
|
+
def __init__(self):
|
|
185
|
+
"""Initialize the rewrite system."""
|
|
186
|
+
self._rules: List[RewriteRule] = []
|
|
187
|
+
self._rule_index: Dict[OpID, List[RewriteRule]] = {}
|
|
188
|
+
self._stats = {
|
|
189
|
+
'total_applications': 0,
|
|
190
|
+
'successful_applications': 0,
|
|
191
|
+
'failed_applications': 0,
|
|
192
|
+
}
|
|
193
|
+
|
|
194
|
+
# Initialize with default rules
|
|
195
|
+
self._initialize_rules()
|
|
196
|
+
|
|
197
|
+
def _initialize_rules(self):
|
|
198
|
+
"""Initialize the default rewrite rules."""
|
|
199
|
+
# Identity rules
|
|
200
|
+
self.add_rule(RewriteRule(
|
|
201
|
+
name="add_zero_left",
|
|
202
|
+
lhs=Pattern.op(OpID.ADD, Pattern.var("x"), Pattern.wildcard()),
|
|
203
|
+
rhs=Pattern.var("x"),
|
|
204
|
+
description="x + 0 -> x"
|
|
205
|
+
))
|
|
206
|
+
|
|
207
|
+
self.add_rule(RewriteRule(
|
|
208
|
+
name="add_zero_right",
|
|
209
|
+
lhs=Pattern.op(OpID.ADD, Pattern.wildcard(), Pattern.var("x")),
|
|
210
|
+
rhs=Pattern.var("x"),
|
|
211
|
+
description="0 + x -> x"
|
|
212
|
+
))
|
|
213
|
+
|
|
214
|
+
self.add_rule(RewriteRule(
|
|
215
|
+
name="mul_one_left",
|
|
216
|
+
lhs=Pattern.op(OpID.MUL, Pattern.var("x"), Pattern.wildcard()),
|
|
217
|
+
rhs=Pattern.var("x"),
|
|
218
|
+
description="x * 1 -> x"
|
|
219
|
+
))
|
|
220
|
+
|
|
221
|
+
self.add_rule(RewriteRule(
|
|
222
|
+
name="mul_one_right",
|
|
223
|
+
lhs=Pattern.op(OpID.MUL, Pattern.wildcard(), Pattern.var("x")),
|
|
224
|
+
rhs=Pattern.var("x"),
|
|
225
|
+
description="1 * x -> x"
|
|
226
|
+
))
|
|
227
|
+
|
|
228
|
+
self.add_rule(RewriteRule(
|
|
229
|
+
name="sub_zero",
|
|
230
|
+
lhs=Pattern.op(OpID.SUB, Pattern.var("x"), Pattern.wildcard()),
|
|
231
|
+
rhs=Pattern.var("x"),
|
|
232
|
+
description="x - 0 -> x"
|
|
233
|
+
))
|
|
234
|
+
|
|
235
|
+
self.add_rule(RewriteRule(
|
|
236
|
+
name="div_one",
|
|
237
|
+
lhs=Pattern.op(OpID.DIV, Pattern.var("x"), Pattern.wildcard()),
|
|
238
|
+
rhs=Pattern.var("x"),
|
|
239
|
+
description="x / 1 -> x"
|
|
240
|
+
))
|
|
241
|
+
|
|
242
|
+
# Self operations
|
|
243
|
+
self.add_rule(RewriteRule(
|
|
244
|
+
name="sub_self",
|
|
245
|
+
lhs=Pattern.op(OpID.SUB, Pattern.var("x"), Pattern.var("x")),
|
|
246
|
+
rhs=Pattern.op(OpID.CONST_ZERO),
|
|
247
|
+
description="x - x -> 0"
|
|
248
|
+
))
|
|
249
|
+
|
|
250
|
+
self.add_rule(RewriteRule(
|
|
251
|
+
name="div_self",
|
|
252
|
+
lhs=Pattern.op(OpID.DIV, Pattern.var("x"), Pattern.var("x")),
|
|
253
|
+
rhs=Pattern.op(OpID.CONST_ONE),
|
|
254
|
+
description="x / x -> 1"
|
|
255
|
+
))
|
|
256
|
+
|
|
257
|
+
# Exponential/Logarithmic
|
|
258
|
+
self.add_rule(RewriteRule(
|
|
259
|
+
name="exp_log",
|
|
260
|
+
lhs=Pattern.op(OpID.EXP, Pattern.op(OpID.LOG, Pattern.var("x"))),
|
|
261
|
+
rhs=Pattern.var("x"),
|
|
262
|
+
description="exp(log(x)) -> x"
|
|
263
|
+
))
|
|
264
|
+
|
|
265
|
+
self.add_rule(RewriteRule(
|
|
266
|
+
name="log_exp",
|
|
267
|
+
lhs=Pattern.op(OpID.LOG, Pattern.op(OpID.EXP, Pattern.var("x"))),
|
|
268
|
+
rhs=Pattern.var("x"),
|
|
269
|
+
description="log(exp(x)) -> x"
|
|
270
|
+
))
|
|
271
|
+
|
|
272
|
+
# Power rules
|
|
273
|
+
self.add_rule(RewriteRule(
|
|
274
|
+
name="power_zero",
|
|
275
|
+
lhs=Pattern.op(OpID.POWER, Pattern.var("x"), Pattern.wildcard()),
|
|
276
|
+
rhs=Pattern.op(OpID.CONST_ONE),
|
|
277
|
+
description="x^0 -> 1"
|
|
278
|
+
))
|
|
279
|
+
|
|
280
|
+
self.add_rule(RewriteRule(
|
|
281
|
+
name="power_one",
|
|
282
|
+
lhs=Pattern.op(OpID.POWER, Pattern.var("x"), Pattern.wildcard()),
|
|
283
|
+
rhs=Pattern.var("x"),
|
|
284
|
+
description="x^1 -> x"
|
|
285
|
+
))
|
|
286
|
+
|
|
287
|
+
self.add_rule(RewriteRule(
|
|
288
|
+
name="power_self",
|
|
289
|
+
lhs=Pattern.op(OpID.POWER, Pattern.var("x"), Pattern.op(OpID.CONST_ONE)),
|
|
290
|
+
rhs=Pattern.var("x"),
|
|
291
|
+
description="x^1 -> x"
|
|
292
|
+
))
|
|
293
|
+
|
|
294
|
+
# Double negation
|
|
295
|
+
self.add_rule(RewriteRule(
|
|
296
|
+
name="neg_neg",
|
|
297
|
+
lhs=Pattern.op(OpID.NEG, Pattern.op(OpID.NEG, Pattern.var("x"))),
|
|
298
|
+
rhs=Pattern.var("x"),
|
|
299
|
+
description="-(-x) -> x"
|
|
300
|
+
))
|
|
301
|
+
|
|
302
|
+
# Inverse rules
|
|
303
|
+
self.add_rule(RewriteRule(
|
|
304
|
+
name="inverse_inverse",
|
|
305
|
+
lhs=Pattern.op(OpID.INVERSE, Pattern.op(OpID.INVERSE, Pattern.var("x"))),
|
|
306
|
+
rhs=Pattern.var("x"),
|
|
307
|
+
description="1/(1/x) -> x"
|
|
308
|
+
))
|
|
309
|
+
|
|
310
|
+
# Square/cube simplifications
|
|
311
|
+
self.add_rule(RewriteRule(
|
|
312
|
+
name="square_neg",
|
|
313
|
+
lhs=Pattern.op(OpID.SQUARE, Pattern.op(OpID.NEG, Pattern.var("x"))),
|
|
314
|
+
rhs=Pattern.op(OpID.SQUARE, Pattern.var("x")),
|
|
315
|
+
description="(-x)^2 -> x^2"
|
|
316
|
+
))
|
|
317
|
+
|
|
318
|
+
self.add_rule(RewriteRule(
|
|
319
|
+
name="cube_neg",
|
|
320
|
+
lhs=Pattern.op(OpID.CUBE, Pattern.op(OpID.NEG, Pattern.var("x"))),
|
|
321
|
+
rhs=Pattern.op(OpID.NEG, Pattern.op(OpID.CUBE, Pattern.var("x"))),
|
|
322
|
+
description="(-x)^3 -> -x^3"
|
|
323
|
+
))
|
|
324
|
+
|
|
325
|
+
# Trigonometric simplifications
|
|
326
|
+
self.add_rule(RewriteRule(
|
|
327
|
+
name="sin_zero",
|
|
328
|
+
lhs=Pattern.op(OpID.SIN, Pattern.op(OpID.CONST_ZERO)),
|
|
329
|
+
rhs=Pattern.op(OpID.CONST_ZERO),
|
|
330
|
+
description="sin(0) -> 0"
|
|
331
|
+
))
|
|
332
|
+
|
|
333
|
+
self.add_rule(RewriteRule(
|
|
334
|
+
name="cos_zero",
|
|
335
|
+
lhs=Pattern.op(OpID.COS, Pattern.op(OpID.CONST_ZERO)),
|
|
336
|
+
rhs=Pattern.op(OpID.CONST_ONE),
|
|
337
|
+
description="cos(0) -> 1"
|
|
338
|
+
))
|
|
339
|
+
|
|
340
|
+
self.add_rule(RewriteRule(
|
|
341
|
+
name="tan_zero",
|
|
342
|
+
lhs=Pattern.op(OpID.TAN, Pattern.op(OpID.CONST_ZERO)),
|
|
343
|
+
rhs=Pattern.op(OpID.CONST_ZERO),
|
|
344
|
+
description="tan(0) -> 0"
|
|
345
|
+
))
|
|
346
|
+
|
|
347
|
+
# Activation function simplifications
|
|
348
|
+
self.add_rule(RewriteRule(
|
|
349
|
+
name="sigmoid_zero",
|
|
350
|
+
lhs=Pattern.op(OpID.SIGMOID, Pattern.op(OpID.CONST_ZERO)),
|
|
351
|
+
rhs=Pattern.op(OpID.CONST_ONE, Pattern.op(OpID.CONST_ZERO)),
|
|
352
|
+
description="sigmoid(0) -> 0.5"
|
|
353
|
+
))
|
|
354
|
+
|
|
355
|
+
self.add_rule(RewriteRule(
|
|
356
|
+
name="relu_positive",
|
|
357
|
+
lhs=Pattern.op(OpID.RELU, Pattern.var("x")),
|
|
358
|
+
rhs=Pattern.var("x"),
|
|
359
|
+
condition=lambda b, h: b.get('x', 0) > 0,
|
|
360
|
+
description="relu(x) -> x for x > 0"
|
|
361
|
+
))
|
|
362
|
+
|
|
363
|
+
self.add_rule(RewriteRule(
|
|
364
|
+
name="relu_negative",
|
|
365
|
+
lhs=Pattern.op(OpID.RELU, Pattern.var("x")),
|
|
366
|
+
rhs=Pattern.op(OpID.CONST_ZERO),
|
|
367
|
+
condition=lambda b, h: b.get('x', 0) <= 0,
|
|
368
|
+
description="relu(x) -> 0 for x <= 0"
|
|
369
|
+
))
|
|
370
|
+
|
|
371
|
+
def add_rule(self, rule: RewriteRule) -> None:
|
|
372
|
+
"""
|
|
373
|
+
Add a rewrite rule to the system.
|
|
374
|
+
|
|
375
|
+
Args:
|
|
376
|
+
rule: The rule to add.
|
|
377
|
+
"""
|
|
378
|
+
self._rules.append(rule)
|
|
379
|
+
|
|
380
|
+
# Index the rule by its LHS operation
|
|
381
|
+
if rule.lhs.op is not None:
|
|
382
|
+
if rule.lhs.op not in self._rule_index:
|
|
383
|
+
self._rule_index[rule.lhs.op] = []
|
|
384
|
+
self._rule_index[rule.lhs.op].append(rule)
|
|
385
|
+
|
|
386
|
+
# Sort rules by priority
|
|
387
|
+
self._rules.sort(key=lambda r: r.priority, reverse=True)
|
|
388
|
+
|
|
389
|
+
def apply_rules(self, handle: int, max_iterations: int = 100) -> int:
|
|
390
|
+
"""
|
|
391
|
+
Apply rewrite rules to an expression until no more rules apply.
|
|
392
|
+
|
|
393
|
+
Args:
|
|
394
|
+
handle: The handle to rewrite.
|
|
395
|
+
max_iterations: Maximum number of iterations.
|
|
396
|
+
|
|
397
|
+
Returns:
|
|
398
|
+
The simplified handle.
|
|
399
|
+
"""
|
|
400
|
+
current = handle
|
|
401
|
+
pool = get_pool()
|
|
402
|
+
|
|
403
|
+
for _ in range(max_iterations):
|
|
404
|
+
applied = False
|
|
405
|
+
|
|
406
|
+
# Get the operation at the root
|
|
407
|
+
node = pool.get_node(current)
|
|
408
|
+
if node is None:
|
|
409
|
+
break
|
|
410
|
+
|
|
411
|
+
# Try to apply rules that match the root operation first
|
|
412
|
+
root_op = node.op
|
|
413
|
+
candidate_rules = self._rule_index.get(root_op, [])
|
|
414
|
+
|
|
415
|
+
# Also consider rules without an LHS op (wildcard patterns)
|
|
416
|
+
for rule in self._rules:
|
|
417
|
+
if rule.lhs.op is None:
|
|
418
|
+
if rule not in candidate_rules:
|
|
419
|
+
candidate_rules.append(rule)
|
|
420
|
+
|
|
421
|
+
# Apply rules in priority order
|
|
422
|
+
for rule in candidate_rules:
|
|
423
|
+
result = rule.apply(current, pool)
|
|
424
|
+
if result is not None and result != current:
|
|
425
|
+
# Rule applied successfully
|
|
426
|
+
current = result
|
|
427
|
+
applied = True
|
|
428
|
+
self._stats['successful_applications'] += 1
|
|
429
|
+
break
|
|
430
|
+
|
|
431
|
+
self._stats['total_applications'] += 1
|
|
432
|
+
|
|
433
|
+
if not applied:
|
|
434
|
+
break
|
|
435
|
+
|
|
436
|
+
return current
|
|
437
|
+
|
|
438
|
+
def normalize(self, handle: int) -> int:
|
|
439
|
+
"""
|
|
440
|
+
Normalize an expression to its simplest form.
|
|
441
|
+
|
|
442
|
+
Args:
|
|
443
|
+
handle: The handle to normalize.
|
|
444
|
+
|
|
445
|
+
Returns:
|
|
446
|
+
The normalized handle.
|
|
447
|
+
"""
|
|
448
|
+
return self.apply_rules(handle)
|
|
449
|
+
|
|
450
|
+
def get_stats(self) -> Dict[str, int]:
|
|
451
|
+
"""Get rewrite system statistics."""
|
|
452
|
+
return {
|
|
453
|
+
'total_rules': len(self._rules),
|
|
454
|
+
'total_applications': self._stats['total_applications'],
|
|
455
|
+
'successful_applications': self._stats['successful_applications'],
|
|
456
|
+
'failed_applications': self._stats['failed_applications'],
|
|
457
|
+
'success_rate': (self._stats['successful_applications'] /
|
|
458
|
+
(self._stats['total_applications'] + 1) * 100),
|
|
459
|
+
}
|
|
460
|
+
|
|
461
|
+
def print_stats(self) -> None:
|
|
462
|
+
"""Print rewrite system statistics."""
|
|
463
|
+
stats = self.get_stats()
|
|
464
|
+
print("=" * 50)
|
|
465
|
+
print("Rewrite System Statistics")
|
|
466
|
+
print("=" * 50)
|
|
467
|
+
print(f"Total Rules: {stats['total_rules']}")
|
|
468
|
+
print(f"Total Applications: {stats['total_applications']}")
|
|
469
|
+
print(f"Successful Applications:{stats['successful_applications']}")
|
|
470
|
+
print(f"Failed Applications: {stats['failed_applications']}")
|
|
471
|
+
print(f"Success Rate: {stats['success_rate']:.2f}%")
|
|
472
|
+
print("=" * 50)
|
|
473
|
+
|
|
474
|
+
|
|
475
|
+
# Global rewrite system
|
|
476
|
+
_REWRITE_SYSTEM: Optional[RewriteSystem] = None
|
|
477
|
+
|
|
478
|
+
|
|
479
|
+
def get_rewrite_system() -> RewriteSystem:
|
|
480
|
+
"""Get or create the global rewrite system."""
|
|
481
|
+
global _REWRITE_SYSTEM
|
|
482
|
+
if _REWRITE_SYSTEM is None:
|
|
483
|
+
_REWRITE_SYSTEM = RewriteSystem()
|
|
484
|
+
return _REWRITE_SYSTEM
|
|
485
|
+
|
|
486
|
+
|
|
487
|
+
def simplify(handle: int) -> int:
|
|
488
|
+
"""
|
|
489
|
+
Simplify an expression using the global rewrite system.
|
|
490
|
+
|
|
491
|
+
Args:
|
|
492
|
+
handle: The handle to simplify.
|
|
493
|
+
|
|
494
|
+
Returns:
|
|
495
|
+
The simplified handle.
|
|
496
|
+
"""
|
|
497
|
+
return get_rewrite_system().normalize(handle)
|