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.
Files changed (74) hide show
  1. metanion-0.1.0/PKG-INFO +57 -0
  2. metanion-0.1.0/README.md +35 -0
  3. metanion-0.1.0/metanion/__init__.py +39 -0
  4. metanion-0.1.0/metanion/algebra/__init__.py +17 -0
  5. metanion-0.1.0/metanion/algebra/rewrite_rules.py +497 -0
  6. metanion-0.1.0/metanion/api/__init__.py +7 -0
  7. metanion-0.1.0/metanion/api/metanion.py +231 -0
  8. metanion-0.1.0/metanion/calculus/__init__.py +16 -0
  9. metanion-0.1.0/metanion/calculus/derivative_rules.py +583 -0
  10. metanion-0.1.0/metanion/calculus/symbolic_differentiator.py +472 -0
  11. metanion-0.1.0/metanion/compile/__init__.py +19 -0
  12. metanion-0.1.0/metanion/compile/bytecode_compiler.py +95 -0
  13. metanion-0.1.0/metanion/compile/lazy_graph.py +355 -0
  14. metanion-0.1.0/metanion/compile/straight_line_program.py +17 -0
  15. metanion-0.1.0/metanion/config.py +117 -0
  16. metanion-0.1.0/metanion/core/__init__.py +12 -0
  17. metanion-0.1.0/metanion/core/dtype_system.py +89 -0
  18. metanion-0.1.0/metanion/core/memory_arena.py +169 -0
  19. metanion-0.1.0/metanion/core/tensor.py +191 -0
  20. metanion-0.1.0/metanion/core/tensor_buffer.py +247 -0
  21. metanion-0.1.0/metanion/core/tensor_shape.py +389 -0
  22. metanion-0.1.0/metanion/data/__init__.py +1 -0
  23. metanion-0.1.0/metanion/data/dataset.py +281 -0
  24. metanion-0.1.0/metanion/data/statistics_injector.py +191 -0
  25. metanion-0.1.0/metanion/exceptions.py +94 -0
  26. metanion-0.1.0/metanion/gp/__init__.py +48 -0
  27. metanion-0.1.0/metanion/gp/bloat_control.py +302 -0
  28. metanion-0.1.0/metanion/gp/crossover.py +275 -0
  29. metanion-0.1.0/metanion/gp/fitness.py +392 -0
  30. metanion-0.1.0/metanion/gp/individual.py +301 -0
  31. metanion-0.1.0/metanion/gp/initialization.py +189 -0
  32. metanion-0.1.0/metanion/gp/mutation.py +261 -0
  33. metanion-0.1.0/metanion/gp/population.py +220 -0
  34. metanion-0.1.0/metanion/gp/safe_ops.py +142 -0
  35. metanion-0.1.0/metanion/gp/selection.py +248 -0
  36. metanion-0.1.0/metanion/io/__init__.py +12 -0
  37. metanion-0.1.0/metanion/io/binary_decoder.py +232 -0
  38. metanion-0.1.0/metanion/io/binary_encoder.py +231 -0
  39. metanion-0.1.0/metanion/io/checkpoint_manager.py +293 -0
  40. metanion-0.1.0/metanion/metanion_engine.py +49 -0
  41. metanion-0.1.0/metanion/model/__init__.py +8 -0
  42. metanion-0.1.0/metanion/model/metanion_layer.py +256 -0
  43. metanion-0.1.0/metanion/model/metanion_model.py +357 -0
  44. metanion-0.1.0/metanion/model/metanion_stack.py +133 -0
  45. metanion-0.1.0/metanion/runtime/__init__.py +24 -0
  46. metanion-0.1.0/metanion/runtime/gc_controller.py +203 -0
  47. metanion-0.1.0/metanion/runtime/jit_cache_manager.py +143 -0
  48. metanion-0.1.0/metanion/symbolic/__init__.py +45 -0
  49. metanion-0.1.0/metanion/symbolic/expression_node.py +91 -0
  50. metanion-0.1.0/metanion/symbolic/handle_utils.py +74 -0
  51. metanion-0.1.0/metanion/symbolic/hash_consing_pool.py +101 -0
  52. metanion-0.1.0/metanion/symbolic/op_enum.py +176 -0
  53. metanion-0.1.0/metanion/symbolic/op_metadata.py +539 -0
  54. metanion-0.1.0/metanion/utils/__init__.py +44 -0
  55. metanion-0.1.0/metanion/utils/cost_model.py +336 -0
  56. metanion-0.1.0/metanion/utils/time_profiler.py +230 -0
  57. metanion-0.1.0/metanion/utils/tree_printer.py +323 -0
  58. metanion-0.1.0/metanion.egg-info/PKG-INFO +57 -0
  59. metanion-0.1.0/metanion.egg-info/SOURCES.txt +72 -0
  60. metanion-0.1.0/metanion.egg-info/dependency_links.txt +1 -0
  61. metanion-0.1.0/metanion.egg-info/requires.txt +1 -0
  62. metanion-0.1.0/metanion.egg-info/top_level.txt +2 -0
  63. metanion-0.1.0/pyproject.toml +3 -0
  64. metanion-0.1.0/setup.cfg +4 -0
  65. metanion-0.1.0/setup.py +20 -0
  66. metanion-0.1.0/tests/__init__.py +4 -0
  67. metanion-0.1.0/tests/run_all_tests.py +70 -0
  68. metanion-0.1.0/tests/test_all_functions.py +508 -0
  69. metanion-0.1.0/tests/test_calculus.py +75 -0
  70. metanion-0.1.0/tests/test_compile.py +69 -0
  71. metanion-0.1.0/tests/test_core.py +72 -0
  72. metanion-0.1.0/tests/test_io.py +54 -0
  73. metanion-0.1.0/tests/test_symbolic.py +80 -0
  74. metanion-0.1.0/tests/test_training.py +100 -0
@@ -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
+ [![PyPI version](https://badge.fury.io/py/metanion.svg)](https://badge.fury.io/py/metanion)
26
+ [![Python 3.8+](https://img.shields.io/badge/python-3.8+-blue.svg)](https://www.python.org/downloads/)
27
+ [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](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())
@@ -0,0 +1,35 @@
1
+ # 🧠 Metanion - Zero-Weight Symbolic Tensor Engine
2
+
3
+ [![PyPI version](https://badge.fury.io/py/metanion.svg)](https://badge.fury.io/py/metanion)
4
+ [![Python 3.8+](https://img.shields.io/badge/python-3.8+-blue.svg)](https://www.python.org/downloads/)
5
+ [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](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)
@@ -0,0 +1,7 @@
1
+ """
2
+ Metanion API - Simple interface for symbolic regression.
3
+ """
4
+
5
+ from .metanion import Metanion
6
+
7
+ __all__ = ['Metanion']