aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorericmarin <maarin.eric@gmail.com>2026-06-20 13:30:33 +0200
committerericmarin <maarin.eric@gmail.com>2026-06-25 20:22:18 +0200
commited622e662cbefc4c16b425666388641e97921ebc (patch)
tree8d525b8740418dc6446d002fc7350c0e3b4da232
parentd85de781b4a825fe0d1217e1e8ae5ec81dd70202 (diff)
downloadvein-master.tar.gz
vein-master.zip
fixed modelHEADmaster
-rw-r--r--README.md2
-rw-r--r--docs/proof.md565
-rw-r--r--examples/double_integrator/double_integrator_pretrain_inv.onnxbin0 -> 12584 bytes
-rw-r--r--examples/pendulum/pendulum_pretrain_con.onnxbin0 -> 6106 bytes
-rw-r--r--vein.py270
-rw-r--r--verify_example.py17
6 files changed, 114 insertions, 740 deletions
diff --git a/README.md b/README.md
index 7989b1c..60222a0 100644
--- a/README.md
+++ b/README.md
@@ -6,8 +6,6 @@ Requires my [fork of Inpla](https://github.com/eric-marin/inpla).
- **make** the executable
- Copy the **inpla** executable in **vein/**
-[Soundness proof](./docs/proof.md)
-
## License
This project is licensed under the GNU Affero General Public License v3.0 - see the [LICENSE](LICENSE) file for details.
diff --git a/docs/proof.md b/docs/proof.md
deleted file mode 100644
index 99c88aa..0000000
--- a/docs/proof.md
+++ /dev/null
@@ -1,565 +0,0 @@
-# Soundness Proof
-## Mathematical Definitions
-![Linear(x, q, r) ~ out => out = q*x + r](https://latex.codecogs.com/svg.image?Linear(x,q,r)\sim&space;out\Rightarrow&space;out=q*x&plus;r&space;)
-![Concrete(k) ~ out => out = ](https://latex.codecogs.com/svg.image?Concrete(k)\sim&space;out\Rightarrow&space;out=k&space;)
-![Add(out, b) ~ a => out = a + b](https://latex.codecogs.com/svg.image?Add(out,b)\sim&space;a\Rightarrow&space;out=a&plus;b&space;)
-![AddCheckLinear(out, x, q, r) ~ b => out = q*x + (r + b)](https://latex.codecogs.com/svg.image?AddCheckLinear(out,x,q,r)\sim&space;b\Rightarrow&space;out=q*x&plus;(r&plus;b))
-![AddCheckConcrete(out, k) ~ b => out = k + b](https://latex.codecogs.com/svg.image?AddCheckConcrete(out,k)\sim&space;b\Rightarrow&space;out=k&plus;b&space;)
-![Mul(out, b) ~ a => out = a * b](https://latex.codecogs.com/svg.image?Mul(out,b)\sim&space;a\Rightarrow&space;out=a*b&space;)
-![MulCheckLinear(out, x, q, r) ~ b => out = q*b*x + r*b](https://latex.codecogs.com/svg.image?MulCheckLinear(out,x,q,r)\sim&space;b\Rightarrow&space;out=q*b*x&plus;r*b&space;)
-![MulCheckConcrete(out, k) ~ b => out = k*b](https://latex.codecogs.com/svg.image?MulCheckConcrete(out,k)\sim&space;b\Rightarrow&space;out=k*b&space;)
-![ReLU(out) ~ x => out = IF (x > 0) THEN x ELSE 0](https://latex.codecogs.com/svg.image?ReLU(out)\sim&space;x\Rightarrow&space;out=\text{IF}\;(x>0)\;\text{THEN}\;x\;\text{ELSE}\;0)
-![Materialize(out) ~ x => out = x](https://latex.codecogs.com/svg.image?Materialize(out)\sim&space;x\Rightarrow&space;out=x&space;)
-
-## Soundness of Translation
-### ReLU
-ONNX ReLU node is defined as:
-![Y = X if X > 0 else 0](https://latex.codecogs.com/svg.image?Y=X\;if\;X>0\;else\;0&space;)
-
-The translation defines the interactions:
-![x_i ~ ReLU(y_i)](https://latex.codecogs.com/svg.image?x_i\sim&space;ReLU(y_i))
-
-By definition this interaction is equal to:
-![y_i = IF (x_i > 0) THEN x_i ELSE 0](https://latex.codecogs.com/svg.image?y_i=\text{IF}\;(x_i>0)\;\text{THEN}\;x_i\;\text{ELSE}\;0)
-
-### Gemm
-ONNX Gemm node is defined as:
-![Y = alpha * A * B + beta * C](https://latex.codecogs.com/svg.image?Y=\alpha\cdot&space;A\cdot&space;B&plus;\beta\cdot&space;C&space;)
-
-The translation defines the interactions:
-![a_i ~ Mul(v_i, Concrete(alpha * b_i))](https://latex.codecogs.com/svg.image?a_i\sim&space;Mul(v_i,Concrete(\alpha*b_i)))
-![Add(...(Add(y_i, v_1), ...), v_n) ~ Concrete(beta * c_i)](https://latex.codecogs.com/svg.image?Add(...(Add(y_i,v_1),...),v_n)\sim&space;Concrete(\beta*c_i))
-
-By definition this interaction is equal to:
-![v_i = alpha * a_i * b_i](https://latex.codecogs.com/svg.image?v_i=\alpha*a_i*b_i&space;)
-![y_i = v_1 + v_2 + ... + v_n + beta * c_i](https://latex.codecogs.com/svg.image?y_i=v_1&plus;v_2&plus;...&plus;v_n&plus;\beta*c_i&space;)
-
-By grouping the operations we get:
-![Y = alpha * A * B + beta * C](https://latex.codecogs.com/svg.image?Y=\alpha\cdot&space;A\cdot&space;B&plus;\beta\cdot&space;C&space;)
-
-### Identiry / Flatten / Reshape / Squeeze / Unsqueeze
-Just identity mapping because wires represent a single element and they are not structured as Tensors.
-![out_i ~ in_i](https://latex.codecogs.com/svg.image?out_i\sim&space;in_i)
-
-### MatMul
-Equal to Gemm with ![alpha=1](https://latex.codecogs.com/svg.image?\inline&space;\alpha=1), ![beta=0](https://latex.codecogs.com/svg.image?\inline&space;\beta=0) and ![C=0](https://latex.codecogs.com/svg.image?\inline&space;&space;C=0).
-
-### Add
-ONNX Add node is defined as:
-![C = A + B](https://latex.codecogs.com/svg.image?&space;C=A&plus;B)
-
-The translation defines the interactions:
-![Add(c_i, b_i) ~ a_i](https://latex.codecogs.com/svg.image?Add(c_i,b_i)\sim&space;a_i)
-
-By definition this interaction is equal to:
-![c_i = a_i + b_i](https://latex.codecogs.com/svg.image?c_i=a_i&plus;b_i)
-
-By grouping the operations we get:
-![C = A + B](https://latex.codecogs.com/svg.image?C=A&plus;B)
-
-### Sub
-ONNX Sub node is defined as:
-![C = A - B](https://latex.codecogs.com/svg.image?C=A-B)
-
-The translation defines the interactions:
-![Add(c_i, neg_b_i) ~ a_i](https://latex.codecogs.com/svg.image?Add(c_i,neg_i)\sim&space;a_i)
-![Mul(neg_b_i, Concrete(-1)) ~ b_i](https://latex.codecogs.com/svg.image?Mul(neg_i,Concrete(-1))\sim&space;b_i)
-
-By definition this interaction is equal to:
-![c_i = a_i + neg_b_i](https://latex.codecogs.com/svg.image?c_i=a_i&plus;neg_i)
-![neg_b_i = -1 * b_i](https://latex.codecogs.com/svg.image?neg_i=-1*b_i)
-
-By grouping the operations we get:
-![C = A - B](https://latex.codecogs.com/svg.image?C=A-B)
-
-### Slice
-ONNX Slice is defined as:
-![out_j = in_{start + (j * step)}](https://latex.codecogs.com/svg.image?&space;out_j=in_{start&plus;(j*step)})
-
-The translations creates a wiring analog to the above definition:
-![out_j ~ in_{start + (j * step)}](https://latex.codecogs.com/svg.image?&space;out_j\sim&space;in_{start&plus;(j*step)})
-
-## Soundness of Interaction Rules
-### Materialize
-The Materialize agent transforms a Linear agent into a tree of explicit mathematical operations
-that are used as final representation for the solver.
-In the Python module the terms are defined as:
-```python
-def TermAdd(a, b):
- return a + b
-def TermMul(a, b):
- return a * b
-def TermReLU(x):
- return z3.If(x > 0, x, 0)
-```
-
-#### Linear >< Materialize
-![Linear(x, q, r) >< Materialize(out) => (1), (2), (3), (4), (5)](https://latex.codecogs.com/svg.image?\inline&space;&space;Linear(x,q,r)><Materialize(out)\Rightarrow&space;(1),(2),(3),(4),(5))
-
-LHS:
-![Linear(x, q, r) ~ wire](https://latex.codecogs.com/svg.image?Linear(x,q,r)\sim&space;wire)
-![Materialize(out) ~ wire](https://latex.codecogs.com/svg.image?Materialize(out)\sim&space;wire)
-![q*x + r = wire](https://latex.codecogs.com/svg.image?q*x&plus;r=wire)
-![out = wire](https://latex.codecogs.com/svg.image?out=wire)
-![out = q*x + r](https://latex.codecogs.com/svg.image?out=q*x&plus;r)
-
-##### Case 1:
-![q = 0 => out ~ Concrete(r), x ~ Eraser](https://latex.codecogs.com/svg.image?q=0\Rightarrow&space;out\sim&space;Concrete(r),x\sim&space;Eraser)
-
-RHS:
-![out = r](https://latex.codecogs.com/svg.image?out=r)
-
-EQUIVALENCE:
-![0*x + r = r => r = r](https://latex.codecogs.com/svg.image?0*x&plus;r=r\Rightarrow&space;r=r)
-
-##### Case 2:
-![q = 1, r = 0 => out ~ x](https://latex.codecogs.com/svg.image?q=1,r=0\Rightarrow&space;out~x)
-
-RHS:
-![x = out](https://latex.codecogs.com/svg.image?x=out)
-![out = x](https://latex.codecogs.com/svg.image?out=x)
-
-EQUIVALENCE:
-![1*x + 0 = x => x = x](https://latex.codecogs.com/svg.image?1*x&plus;0=x\Rightarrow&space;x=x)
-
-##### Case 3:
-![q = 1 => out ~ TermAdd(x, Concrete(r))](https://latex.codecogs.com/svg.image?q=1\Rightarrow&space;out\sim&space;TermAdd(x,Concrete(r)))
-
-RHS:
-![out = x + r](https://latex.codecogs.com/svg.image?out=x&plus;r)
-
-EQUIVALENCE:
-![1*x + r = x + r => x + r = x + r](https://latex.codecogs.com/svg.image?1*x&plus;r=x&plus;r\Rightarrow&space;x&plus;r=x&plus;r)
-
-##### Case 4:
-![r = 0 => out ~ TermMul(Concrete(q), x)](https://latex.codecogs.com/svg.image?r=0\Rightarrow&space;out\sim&space;TermMul(Concrete(q),x))
-
-RHS:
-![out = q*x](https://latex.codecogs.com/svg.image?out=q*x)
-
-EQUIVALENCE:
-![q*x + 0 = q*x => q*x = q*x](https://latex.codecogs.com/svg.image?q*x&plus;0=q*x\Rightarrow&space;q*x=q*x)
-
-##### Case 5:
-![otherwise => out ~ TermAdd(TermMul(Concrete(q), x), Concrete(r))](https://latex.codecogs.com/svg.image?otherwise\Rightarrow&space;out\sim&space;TermAdd(TermMul(Concrete(q),x),Concrete(r)))
-
-RHS:
-![out = q*x + r](https://latex.codecogs.com/svg.image?out=q*x&plus;r)
-
-EQUIVALENCE:
-![q*x + r = q*x + r](https://latex.codecogs.com/svg.image?q*x&plus;r=q*x&plus;r)
-
-#### Concrete >< Materialize
-![Concrete(k) >< Materialize(out) => out ~ Concrete(k)](https://latex.codecogs.com/svg.image?Concrete(k)><Materialize(out)\Rightarrow&space;out\sim&space;Concrete(k))
-
-LHS:
-![Concrete(k) ~ wire](https://latex.codecogs.com/svg.image?Concrete(k)\sim&space;wire)
-![Materialize(out) ~ wire](https://latex.codecogs.com/svg.image?Materialize(out)\sim&space;wire)
-![k = wire](https://latex.codecogs.com/svg.image?k=wire)
-![out = wire](https://latex.codecogs.com/svg.image?out=wire)
-![out = k](https://latex.codecogs.com/svg.image?out=k)
-
-RHS:
-![out = k](https://latex.codecogs.com/svg.image?out=k)
-
-EQUIVALENCE:
-![k = k](https://latex.codecogs.com/svg.image?k=k)
-
-### Add
-#### Linear >< Add
-![Linear(x, q, r) >< Add(out, b) => b ~ AddCheckLinear(out, x, q, r)](https://latex.codecogs.com/svg.image?Linear(x,q,r)><Add(out,b)\Rightarrow&space;b\sim&space;AddCheckLinear(out,x,q,r))
-
-LHS:
-![Linear(x, q, r) ~ wire](https://latex.codecogs.com/svg.image?Linear(x,q,r)\sim&space;wire)
-![Add(out, b) ~ wire](https://latex.codecogs.com/svg.image?Add(out,b)\sim&space;wire)
-![q*x + r = wire](https://latex.codecogs.com/svg.image?q*x&plus;r=wire)
-![out = wire + b](https://latex.codecogs.com/svg.image?out=wire&plus;b)
-![out = q*x + r + b](https://latex.codecogs.com/svg.image?out=q*x&plus;r&plus;b)
-
-RHS:
-![out = q*x + (r + b)](https://latex.codecogs.com/svg.image?out=q*x&plus;(r&plus;b))
-
-EQUIVALENCE:
-![q*x + r + b = q*x + (r + b) => q*x + (r + b) = q*x + (r + b)](https://latex.codecogs.com/svg.image?q*x&plus;r&plus;b=q*x&plus;(r&plus;b)\Rightarrow&space;q*x&plus;(r&plus;b)=q*x&plus;(r&plus;b))
-
-#### Concrete >< Add
-![Concrete(k) >< Add(out, b) => (1), (2)](https://latex.codecogs.com/svg.image?Concrete(k)><Add(out,b)\Rightarrow&space;(1),(2))
-
-LHS:
-![Concrete(k) ~ wire](https://latex.codecogs.com/svg.image?Concrete(k)\sim&space;wire)
-![Add(out, b) ~ wire](https://latex.codecogs.com/svg.image?Add(out,b)\sim&space;wire)
-![k = wire](https://latex.codecogs.com/svg.image?k=wire)
-![out = wire + b](https://latex.codecogs.com/svg.image?out=wire&plus;b)
-![out = k + b](https://latex.codecogs.com/svg.image?out=k&plus;b)
-
-##### Case 1:
-![k = 0 => out ~ b](https://latex.codecogs.com/svg.image?k=0\Rightarrow&space;out\sim&space;b)
-
-RHS:
-![out = b](https://latex.codecogs.com/svg.image?out=b)
-
-EQUIVALENCE:
-![0 + b = b => b = b](https://latex.codecogs.com/svg.image?0&plus;b=b\Rightarrow&space;b=b)
-
-##### Case 2:
-![otherwise => b ~ AddCheckConcrete(out, k)](https://latex.codecogs.com/svg.image?otherwise\Rightarrow&space;b\sim&space;AddCheckConcrete(out,k))
-
-RHS:
-![out = k + b](https://latex.codecogs.com/svg.image?out=k&plus;b)
-
-EQUIVALENCE:
-![k + b = k + b](https://latex.codecogs.com/svg.image?k&plus;b=k&plus;b)
-
-#### Linear >< AddCheckLinear
-![Linear(y, s, t) >< AddCheckLinear(out, x, q, r) => (1), (2), (3), (4)](https://latex.codecogs.com/svg.image?Linear(y,s,t)><AddCheckLinear(out,x,q,r)\Rightarrow&space;(1),(2),(3),(4))
-
-LHS:
-![Linear(y, s, t) ~ wire](https://latex.codecogs.com/svg.image?Linear(y,s,t)\sim&space;wire)
-![AddCheckLinear(out, x, q, r) ~ wire](https://latex.codecogs.com/svg.image?AddCheckLinear(out,x,q,r)\sim&space;wire)
-![s*y + t = wire](https://latex.codecogs.com/svg.image?s*y&plus;t=wire)
-![out = q*x + (r + wire)](https://latex.codecogs.com/svg.image?out=q*x&plus;(r&plus;wire))
-![out = q*x + (r + s*y + t)](https://latex.codecogs.com/svg.image?out=q*x&plus;(r&plus;s*y&plus;t))
-
-##### Case 1:
-![q,r,s,t = 0 => out ~ Concrete(0), x ~ Eraser, y ~ Eraser](https://latex.codecogs.com/svg.image?q,r,s,t=0\Rightarrow&space;out\sim&space;Concrete(0),x\sim&space;Eraser,y\sim&space;Eraser)
-
-RHS:
-![out = 0](https://latex.codecogs.com/svg.image?out=0)
-
-EQUIVALENCE:
-![0*x + (0 + 0*y + 0) = 0 => 0 = 0](https://latex.codecogs.com/svg.image?0*x&plus;(0&plus;0*y&plus;0)=0\Rightarrow&space;0=0)
-
-##### Case 2:
-![s,t = 0 => out ~ Linear(x, q, r), y ~ Eraser](https://latex.codecogs.com/svg.image?s,t=0\Rightarrow&space;out\sim&space;Linear(x,q,r),y\sim&space;Eraser)
-
-RHS:
-![out = q*x + r](https://latex.codecogs.com/svg.image?out=q*x&plus;r)
-
-EQUIVALENCE:
-![q*x + (r + 0*y + 0) = q*x + r => q*x + r = q*x + r](https://latex.codecogs.com/svg.image?q*x&plus;(r&plus;0*y&plus;0)=q*x&plus;r\Rightarrow&space;q*x&plus;r=q*x&plus;r)
-
-##### Case 3:
-![q, r = 0 => out ~ Linear(y, s, t), x ~ Eraser](https://latex.codecogs.com/svg.image?q,r=0\Rightarrow&space;out\sim&space;Linear(y,s,t),x\sim&space;Eraser)
-
-RHS:
-![out = s*y + t](https://latex.codecogs.com/svg.image?out=s*y&plus;t)
-
-EQUIVALENCE:
-![0*x + (0 + s*y + t) = s*y + t => s*y + t = s*y + t](https://latex.codecogs.com/svg.image?0*x&plus;(0&plus;s*y&plus;t)=s*y&plus;t\Rightarrow&space;s*y&plus;t=s*y&plus;t)
-
-##### Case 4:
-![otherwise => Linear(x, q, r) ~ Materialize(out_x), Linear(y, s, t) ~ Materialize(out_y), out ~ Linear(TermAdd(out_x, out_y), 1, 0)](https://latex.codecogs.com/svg.image?otherwise\Rightarrow&space;Linear(x,q,r)\sim&space;Materialize(out_x),Linear(y,s,t)\sim&space;Materialize(out_y),out\sim&space;Linear(TermAdd(out_x,out_y),1,0))
-
-RHS:
-![Linear(x, q, r) ~ wire_1](https://latex.codecogs.com/svg.image?Linear(x,q,r)\sim&space;wire_1)
-![Materialize(out_x) ~ wire_1](https://latex.codecogs.com/svg.image?Materialize(out_x)\sim&space;wire_1)
-![q*x + r = wire_1](https://latex.codecogs.com/svg.image?q*x&plus;r=wire_1)
-![out_x = wire_1](https://latex.codecogs.com/svg.image?out_x=wire_1)
-![Linear(y, s, t) ~ wire_2](https://latex.codecogs.com/svg.image?Linear(y,s,t)\sim&space;wire_2)
-![Materialize(out_y) ~ wire_2](https://latex.codecogs.com/svg.image?Materialize(out_y)\sim&space;wire_2)
-![s*y + t = wire_2](https://latex.codecogs.com/svg.image?s*y&plus;t=wire_2)
-![out_y = wire_2](https://latex.codecogs.com/svg.image?out_y=wire_2)
-![out = 1*TermAdd(out_x, out_y) + 0](https://latex.codecogs.com/svg.image?out=1*TermAdd(out_x,out_y)&plus;0)
-Because `TermAdd(a, b)` is defined as `a+b`:
-![out = 1*(q*x + r + s*y + t) + 0](https://latex.codecogs.com/svg.image?out=1*(q*x&plus;r&plus;s*y&plus;t)&plus;0)
-
-EQUIVALENCE:
-![q*x + (r + s*y + t) = 1*(q*x + r + s*y + t) + 0 => q*x + r + s*y + t = q*x + r + s*y + t](https://latex.codecogs.com/svg.image?q*x&plus;(r&plus;s*y&plus;t)=1*(q*x&plus;r&plus;s*y&plus;t)&plus;0\Rightarrow&space;q*x&plus;r&plus;s*y&plus;t=q*x&plus;r&plus;s*y&plus;t)
-
-#### Concrete >< AddCheckLinear
-![Concrete(j) >< AddCheckLinear(out, x, q, r) => out ~ Linear(x, q, r + j)](https://latex.codecogs.com/svg.image?Concrete(j)><AddCheckLinear(out,x,q,r)\Rightarrow&space;out\sim&space;Linear(x,q,r&plus;j))
-
-LHS:
-![Concrete(j) ~ wire](https://latex.codecogs.com/svg.image?Concrete(j)\sim&space;wire)
-![AddCheckLinear(out, x, q, r) ~ wire](https://latex.codecogs.com/svg.image?AddCheckLinear(out,x,q,r)\sim&space;wire)
-![j = wire](https://latex.codecogs.com/svg.image?j=wire)
-![out = q*x + (r + wire)](https://latex.codecogs.com/svg.image?out=q*x&plus;(r&plus;wire))
-![out = q*x + (r + j)](https://latex.codecogs.com/svg.image?out=q*x&plus;(r&plus;j))
-
-RHS:
-![out = q*x + (r + j)](https://latex.codecogs.com/svg.image?out=q*x&plus;(r&plus;j))
-
-EQUIVALENCE:
-![q*x + (r + j) = q*x + (r + j)](https://latex.codecogs.com/svg.image?q*x&plus;(r&plus;j)=q*x&plus;(r&plus;j))
-
-#### Linear >< AddCheckConcrete
-![Linear(y, s, t) >< AddCheckConcrete(out, k) => out ~ Linear(y, s, t + k)](https://latex.codecogs.com/svg.image?Linear(y,s,t)><AddCheckConcrete(out,k)\Rightarrow&space;out\sim&space;Linear(y,s,t&plus;k))
-
-LHS:
-![Linear(y, s, t) ~ wire](https://latex.codecogs.com/svg.image?Linear(y,s,t)\sim&space;wire)
-![AddCheckConcrete(out, k) ~ wire](https://latex.codecogs.com/svg.image?AddCheckConcrete(out,k)\sim&space;wire)
-![s*y + t = wire](https://latex.codecogs.com/svg.image?s*y&plus;t=wire)
-![out = k + wire](https://latex.codecogs.com/svg.image?out=k&plus;wire)
-![out = k + s*y + t](https://latex.codecogs.com/svg.image?out=k&plus;s*y&plus;t)
-
-RHS:
-![out = s*y + (t + k)](https://latex.codecogs.com/svg.image?out=s*y&plus;(t&plus;k))
-
-EQUIVALENCE:
-![k + s*y + t = s*y + (t + k) => s*y + (t + k) = s*y + (t + k)](https://latex.codecogs.com/svg.image?k&plus;s*y&plus;t=s*y&plus;(t&plus;k)\Rightarrow&space;s*y&plus;(t&plus;k)=s*y&plus;(t&plus;k))
-
-#### Concrete >< AddCheckConcrete
-![Concrete(j) >< AddCheckConcrete(out, k) => (1), (2)](https://latex.codecogs.com/svg.image?Concrete(j)><AddCheckConcrete(out,k)\Rightarrow&space;(1),(2))
-
-LHS:
-![Concrete(j) ~ wire](https://latex.codecogs.com/svg.image?Concrete(j)\sim&space;wire)
-![AddCheckConcrete(out, k) ~ wire](https://latex.codecogs.com/svg.image?AddCheckConcrete(out,k)\sim&space;wire)
-![j = wire](https://latex.codecogs.com/svg.image?j=wire)
-![out = k + wire](https://latex.codecogs.com/svg.image?out=k&plus;wire)
-![out = k + j](https://latex.codecogs.com/svg.image?out=k&plus;j)
-
-##### Case 1:
-![j = 0 => out ~ Concrete(k)](https://latex.codecogs.com/svg.image?j=0\Rightarrow&space;out\sim&space;Concrete(k))
-
-RHS:
-![out = k](https://latex.codecogs.com/svg.image?out=k)
-
-EQUIVALENCE:
-![k + 0 = k => k = k](https://latex.codecogs.com/svg.image?k&plus;0=k\Rightarrow&space;k=k)
-
-##### Case 2:
-![otherwise => out ~ Concrete(k + j)](https://latex.codecogs.com/svg.image?otherwise\Rightarrow&space;out\sim&space;Concrete(k&plus;j))
-
-RHS:
-![out = k + j](https://latex.codecogs.com/svg.image?out=k&plus;j)
-
-EQUIVALENCE:
-![k + j = k + j](https://latex.codecogs.com/svg.image?k&plus;j=k&plus;j)
-
-### Mul
-#### Linear >< Mul
-![Linear(x, q, r) >< Mul(out, b) => b ~ MulCheckLinear(out, x, q, r)](https://latex.codecogs.com/svg.image?Linear(x,q,r)><Mul(out,b)\Rightarrow&space;b\sim&space;MulCheckLinear(out,x,q,r))
-
-LHS:
-![Linear(x, q, r) ~ wire](https://latex.codecogs.com/svg.image?Linear(x,q,r)\sim&space;wire)
-![Mul(out, b) ~ wire](https://latex.codecogs.com/svg.image?Mul(out,b)\sim&space;wire)
-![q*x + r = wire](https://latex.codecogs.com/svg.image?q*x&plus;r=wire)
-![out = wire * b](https://latex.codecogs.com/svg.image?out=wire*b)
-![out = (q*x + r) * b](https://latex.codecogs.com/svg.image?out=(q*x&plus;r)*b)
-
-RHS:
-![out = q*b*x + r*b](https://latex.codecogs.com/svg.image?out=q*b*x&plus;r*b)
-
-EQUIVALENCE:
-![(q*x + r) * b = q*b*x + r*b => q*b*x + r*b = q*b*x + r*b](https://latex.codecogs.com/svg.image?(q*x&plus;r)*b=q*b*x&plus;r*b\Rightarrow&space;q*b*x&plus;r*b=q*b*x&plus;r*b)
-
-#### Concrete >< Mul
-![Concrete(k) >< Mul(out, b) => (1), (2), (3)](https://latex.codecogs.com/svg.image?Concrete(k)><Mul(out,b)\Rightarrow&space;(1),(2),(3))
-
-LHS:
-![Concrete(k) ~ wire](https://latex.codecogs.com/svg.image?Concrete(k)\sim&space;wire)
-![Mul(out, b) ~ wire](https://latex.codecogs.com/svg.image?Mul(out,b)\sim&space;wire)
-![k = wire](https://latex.codecogs.com/svg.image?k=wire)
-![out = wire * b](https://latex.codecogs.com/svg.image?out=wire*b)
-![out = k * b](https://latex.codecogs.com/svg.image?out=k*b)
-
-##### Case 1:
-![k = 0 => out ~ Concrete(0), b ~ Eraser](https://latex.codecogs.com/svg.image?k=0\Rightarrow&space;out\sim&space;Concrete(0),b\sim&space;Eraser)
-
-RHS:
-![out = 0](https://latex.codecogs.com/svg.image?out=0)
-
-EQUIVALENCE:
-![0 * b = 0 => 0 = 0](https://latex.codecogs.com/svg.image?0*b=0\Rightarrow&space;0=0)
-
-##### Case 2:
-![k = 1 => out ~ b](https://latex.codecogs.com/svg.image?k=1\Rightarrow&space;out\sim&space;b)
-
-RHS:
-![out = b](https://latex.codecogs.com/svg.image?out=b)
-
-EQUIVALENCE:
-![1 * b = b => b = b](https://latex.codecogs.com/svg.image?1*b=b\Rightarrow&space;b=b)
-
-##### Case 3:
-![otherwise => b ~ MulCheckConcrete(out, k)](https://latex.codecogs.com/svg.image?otherwise\Rightarrow&space;b\sim&space;MulCheckConcrete(out,k))
-
-RHS:
-![out = k * b](https://latex.codecogs.com/svg.image?out=k*b)
-
-EQUIVALENCE:
-![k * b = k * b](https://latex.codecogs.com/svg.image?k*b=k*b)
-
-#### Linear >< MulCheckLinear
-![Linear(y, s, t) >< MulCheckLinear(out, x, q, r) => (1), (2)](https://latex.codecogs.com/svg.image?Linear(y,s,t)><MulCheckLinear(out,x,q,r)\Rightarrow&space;(1),(2))
-
-LHS:
-![Linear(y, s, t) ~ wire](https://latex.codecogs.com/svg.image?Linear(y,s,t)\sim&space;wire)
-![MulCheckLinear(out, x, q, r) ~ wire](https://latex.codecogs.com/svg.image?MulCheckLinear(out,x,q,r)\sim&space;wire)
-![s*y + t = wire](https://latex.codecogs.com/svg.image?s*y&plus;t=wire)
-![out = q*wire*x + r*wire](https://latex.codecogs.com/svg.image?out=q*wire*x&plus;r*wire)
-![out = q*(s*y + t)*x + r*(s*y + t)](https://latex.codecogs.com/svg.image?out=q*(s*y&plus;t)*x&plus;r*(s*y&plus;t))
-
-##### Case 1:
-![(q,r = 0) or (s,t = 0) => x ~ Eraser, y ~ Eraser, out ~ Concrete(0)](https://latex.codecogs.com/svg.image?(q,r=0)\lor(s,t=0)\Rightarrow&space;x\sim&space;Eraser,y\sim&space;Eraser,out\sim&space;Concrete(0))
-
-RHS:
-![out = 0](https://latex.codecogs.com/svg.image?out=0)
-
-EQUIVALENCE:
-![0*(s*y + t)*x + 0*(s*y + t) = 0 => 0 = 0](https://latex.codecogs.com/svg.image?0*(s*y&plus;t)*x&plus;0*(s*y&plus;t)=0\Rightarrow&space;0=0)
-![or](https://latex.codecogs.com/svg.image?\lor)
-![q*(0*y + 0)*x + r*(0*y + 0) = 0 => 0 = 0](https://latex.codecogs.com/svg.image?q*(0*y&plus;0)*x&plus;r*(0*y&plus;0)=0\Rightarrow&space;0=0)
-
-##### Case 2:
-![otherwise => Linear(x, q, r) ~ Materialize(out_x), Linear(y, s, t) ~ Materialize(out_y), out ~ Linear(TermMul(out_x, out_y), 1, 0)](https://latex.codecogs.com/svg.image?otherwise\Rightarrow&space;Linear(x,q,r)\sim&space;Materialize(out_x),Linear(y,s,t)\sim&space;Materialize(out_y),out\sim&space;Linear(TermMul(out_x,out_y),1,0))
-
-RHS:
-![Linear(x, q, r) ~ wire_1](https://latex.codecogs.com/svg.image?Linear(x,q,r)\sim&space;wire_1)
-![Materialize(out_x) ~ wire_1](https://latex.codecogs.com/svg.image?Materialize(out_x)\sim&space;wire_1)
-![q*x + r = wire_1](https://latex.codecogs.com/svg.image?q*x&plus;r=wire_1)
-![out_x = wire_1](https://latex.codecogs.com/svg.image?out_x=wire_1)
-![Linear(y, s, t) ~ wire_2](https://latex.codecogs.com/svg.image?Linear(y,s,t)\sim&space;wire_2)
-![Materialize(out_y) ~ wire_2](https://latex.codecogs.com/svg.image?Materialize(out_y)\sim&space;wire_2)
-![s*y + t = wire_2](https://latex.codecogs.com/svg.image?s*y&plus;t=wire_2)
-![out_y = wire_2](https://latex.codecogs.com/svg.image?out_y=wire_2)
-![out = 1*TermMul(out_x, out_y) + 0](https://latex.codecogs.com/svg.image?out=1*TermMul(out_x,out_y)&plus;0)
-Because `TermMul(a, b)` is defined as `a*b`:
-![out = 1*(q*x + r)*(s*y + t) + 0](https://latex.codecogs.com/svg.image?out=1*(q*x&plus;r)*(s*y&plus;t)&plus;0)
-
-EQUIVALENCE:
-![q*(s*y + t)*x + r*(s*y + t) = 1*(q*x + r)*(s*y + t) =>
-q*(s*y + t)*x + r*(s*y + t) = (q*x + r)*(s*y + t) =>
-q*(s*y + t)*x + r*(s*y + t) = q*(s*y + t)*x + r*(s*y + t)](https://latex.codecogs.com/svg.image?q*(s*y&plus;t)*x&plus;r*(s*y&plus;t)=1*(q*x&plus;r)*(s*y&plus;t)\Rightarrow&space;q*(s*y&plus;t)*x&plus;r*(s*y&plus;t)=(q*x&plus;r)*(s*y&plus;t)\Rightarrow&space;q*(s*y&plus;t)*x&plus;r*(s*y&plus;t)=q*(s*y&plus;t)*x&plus;r*(s*y&plus;t))
-
-
-#### Concrete >< MulCheckLinear
-![Concrete(j) >< MulCheckLinear(out, x, q, r) => out ~ Linear(x, q * j, r * j)](https://latex.codecogs.com/svg.image?Concrete(j)><MulCheckLinear(out,x,q,r)\Rightarrow&space;out\sim&space;Linear(x,q*j,r*j))
-
-LHS:
-![Concrete(j) ~ wire](https://latex.codecogs.com/svg.image?Concrete(j)\sim&space;wire)
-![MulCheckLinear(out, x, q, r) ~ wire](https://latex.codecogs.com/svg.image?MulCheckLinear(out,x,q,r)\sim&space;wire)
-![j = wire](https://latex.codecogs.com/svg.image?j=wire)
-![out = q*wire*x + r*wire](https://latex.codecogs.com/svg.image?out=q*wire*x&plus;r*wire)
-![out = q*j*x + r*j](https://latex.codecogs.com/svg.image?out=q*j*x&plus;r*j)
-
-RHS:
-![out = q*j*x + r*j](https://latex.codecogs.com/svg.image?out=q*j*x&plus;r*j)
-
-EQUIVALENCE:
-![q*j*x + r*j = q*j*x + r*j](https://latex.codecogs.com/svg.image?q*j*x&plus;r*j=q*j*x&plus;r*j)
-
-#### Linear >< MulCheckConcrete
-![Linear(y, s, t) >< MulCheckConcrete(out, k) => out ~ Linear(y, s * k, t * k)](https://latex.codecogs.com/svg.image?Linear(y,s,t)><MulCheckConcrete(out,k)\Rightarrow&space;out\sim&space;Linear(y,s*k,t*k))
-
-LHS:
-![Linear(y, s, t) ~ wire](https://latex.codecogs.com/svg.image?Linear(y,s,t)\sim&space;wire)
-![MulCheckConcrete(out, k) ~ wire](https://latex.codecogs.com/svg.image?MulCheckConcrete(out,k)\sim&space;wire)
-![s*y + t = wire](https://latex.codecogs.com/svg.image?s*y&plus;t=wire)
-![out = k * wire](https://latex.codecogs.com/svg.image?out=k*wire)
-![out = k * (s*y + t)](https://latex.codecogs.com/svg.image?out=k*(s*y&plus;t))
-
-RHS:
-![out = s*k*y + t*k](https://latex.codecogs.com/svg.image?out=s*k*y&plus;t*k)
-
-EQUIVALENCE:
-![k * (s*y + t) = s*k*y + t*k => s*k*y + t*k = s*k*y + t*k](https://latex.codecogs.com/svg.image?k*(s*y&plus;t)=s*k*y&plus;t*k\Rightarrow&space;s*k*y&plus;t*k=s*k*y&plus;t*k)
-
-
-#### Concrete >< MulCheckConcrete
-![Concrete(j) >< MulCheckConcrete(out, k) => (1), (2), (3)](https://latex.codecogs.com/svg.image?Concrete(j)><MulCheckConcrete(out,k)\Rightarrow&space;(1),(2),(3))
-
-LHS:
-![Concrete(j) ~ wire](https://latex.codecogs.com/svg.image?Concrete(j)\sim&space;wire)
-![MulCheckConcrete(out, k) ~ wire](https://latex.codecogs.com/svg.image?MulCheckConcrete(out,k)\sim&space;wire)
-![j = wire](https://latex.codecogs.com/svg.image?j=wire)
-![out = k * wire](https://latex.codecogs.com/svg.image?out=k*wire)
-![out = k * j](https://latex.codecogs.com/svg.image?out=k*j)
-
-##### Case 1:
-![j = 0 => out ~ Concrete(0)](https://latex.codecogs.com/svg.image?j=0\Rightarrow&space;out\sim&space;Concrete(0))
-
-RHS:
-![out = 0](https://latex.codecogs.com/svg.image?out=0)
-
-EQUIVALENCE:
-![k * 0 = 0 => 0 = 0](https://latex.codecogs.com/svg.image?k*0=0\Rightarrow&space;0=0)
-
-##### Case 2:
-![j = 1 => out ~ Concrete(k)](https://latex.codecogs.com/svg.image?j=1\Rightarrow&space;out\sim&space;Concrete(k))
-
-RHS:
-![out = k](https://latex.codecogs.com/svg.image?out=k)
-
-EQUIVALENCE:
-![k * 1 = k => k = k](https://latex.codecogs.com/svg.image?k*1=k\Rightarrow&space;k=k)
-
-##### Case 3:
-![otherwise => out ~ Concrete(k * j)](https://latex.codecogs.com/svg.image?otherwise\Rightarrow&space;out\sim&space;Concrete(k*j))
-
-RHS:
-![out = k * j](https://latex.codecogs.com/svg.image?out=k*j)
-
-EQUIVALENCE:
-![k * j = k * j](https://latex.codecogs.com/svg.image?k*j=k*j)
-
-### ReLU
-#### Linear >< ReLU
-![Linear(x, q, r) >< ReLU(out) => Linear(x, q, r) ~ Materialize(out_x), out ~ Linear(TermReLU(out_x), 1, 0)](https://latex.codecogs.com/svg.image?Linear(x,q,r)><ReLU(out)\Rightarrow&space;Linear(x,q,r)\sim&space;Materialize(out_x),out\sim&space;Linear(TermReLU(out_x),1,0))
-
-LHS:
-![Linear(x, q, r) ~ wire](https://latex.codecogs.com/svg.image?Linear(x,q,r)\sim&space;wire)
-![ReLU(out) ~ wire](https://latex.codecogs.com/svg.image?ReLU(out)\sim&space;wire)
-![q*x + r = wire](https://latex.codecogs.com/svg.image?q*x&plus;r=wire)
-![out = IF wire > 0 THEN wire ELSE 0](https://latex.codecogs.com/svg.image?out=\text{IF}\;wire>0\;\text{THEN}\;wire\;\text{ELSE}\;0)
-![out = IF (q*x + r) > 0 THEN (q*x + r) ELSE 0](https://latex.codecogs.com/svg.image?out=\text{IF}\;(q*x&plus;r)>0\;\text{THEN}\;(q*x&plus;r)\;\text{ELSE}\;0)
-
-RHS:
-![Linear(x, q, r) ~ wire](https://latex.codecogs.com/svg.image?Linear(x,q,r)\sim&space;wire)
-![Materialize(out_x) ~ wire](https://latex.codecogs.com/svg.image?Materialize(out_x)\sim&space;wire)
-![q*x + r = wire](https://latex.codecogs.com/svg.image?q*x&plus;r=wire)
-![out_x = wire](https://latex.codecogs.com/svg.image?out_x=wire)
-![out = 1*TermReLU(out_x) + 0](https://latex.codecogs.com/svg.image?out=1*TermReLU(out_x)&plus;0)
-Because `TermReLU(x)` is defined as `z3.If(x > 0, x, 0)`:
-![out = 1*(IF (q*x + r) > 0 THEN (q*x + r) ELSE 0) + 0](https://latex.codecogs.com/svg.image?out=1*(\text{IF}\;(q*x&plus;r)>0\;\text{THEN}\;(q*x&plus;r)\;\text{ELSE}\;0)&plus;0)
-
-EQUIVALENCE:
-![IF (q*x + r) > 0 THEN (q*x + r) ELSE 0 = 1*(IF (q*x + r) > 0 THEN (q*x + r) ELSE 0) + 0 =>
-IF (q*x + r) > 0 THEN (q*x + r) ELSE 0 = IF (q*x + r) > 0 THEN (q*x + r) ELSE 0](https://latex.codecogs.com/svg.image?\text{IF}\;(q*x&plus;r)>0\;\text{THEN}\;(q*x&plus;r)\;\text{ELSE}\;0=1*(\text{IF}\;(q*x&plus;r)>0\;\text{THEN}\;(q*x&plus;r)\;\text{ELSE}\;0)&plus;0\Rightarrow\text{IF}\;(q*x&plus;r)>0\;\text{THEN}\;(q*x&plus;r)\;\text{ELSE}\;0=\text{IF}\;(q*x&plus;r)>0\;\text{THEN}\;(q*x&plus;r)\;\text{ELSE}\;0)
-
-
-#### Concrete >< ReLU
-![Concrete(k) >< ReLU(out) => (1), (2)](https://latex.codecogs.com/svg.image?Concrete(k)><ReLU(out)\Rightarrow&space;(1),(2))
-
-LHS:
-![Concrete(k) ~ wire](https://latex.codecogs.com/svg.image?Concrete(k)\sim&space;wire)
-![ReLU(out) ~ wire](https://latex.codecogs.com/svg.image?ReLU(out)\sim&space;wire)
-![k = wire](https://latex.codecogs.com/svg.image?k=wire)
-![out = IF wire > 0 THEN wire ELSE 0](https://latex.codecogs.com/svg.image?out=\text{IF}\;wire>0\;\text{THEN}\;wire\;\text{ELSE}\;0)
-![out = IF k > 0 THEN k ELSE 0](https://latex.codecogs.com/svg.image?out=\text{IF}\;k>0\;\text{THEN}\;k\;\text{ELSE}\;0)
-
-##### Case 1:
-![k > 0 => out ~ Concrete(k)](https://latex.codecogs.com/svg.image?k>0\Rightarrow&space;out\sim&space;Concrete(k))
-
-RHS:
-![out = k](https://latex.codecogs.com/svg.image?out=k)
-
-EQUIVALENCE:
-![IF true THEN k ELSE 0 = k => k = k](https://latex.codecogs.com/svg.image?\text{IF}\;true\;\text{THEN}\;k\;\text{ELSE}\;0=k\Rightarrow&space;k=k)
-
-##### Case 2:
-![k <= 0 => out ~ Concrete(0)](https://latex.codecogs.com/svg.image?k\leq&space;0\Rightarrow&space;out\sim&space;Concrete(0))
-
-RHS:
-![out = 0](https://latex.codecogs.com/svg.image?out=0)
-
-EQUIVALENCE:
-![IF false THEN k ELSE 0 = 0 => 0 = 0](https://latex.codecogs.com/svg.image?\text{IF}\;false\;\text{THEN}\;k\;\text{ELSE}\;0=0\Rightarrow&space;0=0)
-
-## Soundness of Reduction
-Let ![IN_0](https://latex.codecogs.com/svg.image?\mathrm{IN_0}) be the Interaction Net translated from a Neural Network ![NN](https://latex.codecogs.com/svg.image?\inline&space;\mathrm{NN}). Let ![IN_n](https://latex.codecogs.com/svg.image?\inline&space;\mathrm{IN_n}) be the state of the net
-after ![n](https://latex.codecogs.com/svg.image?\inline&space;n) reduction steps. Then ![forall n in N, [IN_n] = [NN]](https://latex.codecogs.com/svg.image?\inline&space;\forall&space;n\in\mathbb{N},[\mathrm{IN_n}]=[\mathrm{NN}]).
-
-### Proof by Induction
-- Base Case (![n = 0](https://latex.codecogs.com/svg.image?\inline&space;n=0)): By the [Soundness of Translation](#soundness-of-translation), the initial net ![IN_0](https://latex.codecogs.com/svg.image?\mathrm{IN_0}) is constructed such that
-its semantics ![[IN_0]](https://latex.codecogs.com/svg.image?\inline&space;[\mathrm{IN_0}]) exactly match the mathematical definition of the ONNX nodes in ![NN](https://latex.codecogs.com/svg.image?\inline&space;\mathrm{NN}).
-- Induction Step (![n -> n + 1](https://latex.codecogs.com/svg.image?\inline&space;n\to&space;n&plus;1)): Assume ![[IN_n] = [NN]](https://latex.codecogs.com/svg.image?\inline&space;[\mathrm{IN_n}]=[\mathrm{NN}]). If ![IN_n](https://latex.codecogs.com/svg.image?\inline&space;\mathrm{IN_n}) is in normal form, the proof is complete.
-Otherwise, there exists an active pair ![A](https://latex.codecogs.com/svg.image?\inline&space;A) that reduces ![IN_n to IN_{n+1}](https://latex.codecogs.com/svg.image?\inline&space;\mathrm{IN_n}\Rightarrow&space;\mathrm{IN_{n&plus;1}}).
-By the [Soundness of Interaction Rules](#soundness-of-interaction-rules), the mathematical definition is preserved after any reduction step,
-it follows that ![[IN_{n+1}] = [IN_n]](https://latex.codecogs.com/svg.image?\inline&space;[\mathrm{IN_{n&plus;1}}]=[\mathrm{IN_n}]). By the inductive hypothesis, ![[IN_{n+1}] = [NN]](https://latex.codecogs.com/svg.image?\inline&space;[\mathrm{IN_{n&plus;1}}]=[\mathrm{NN}]).
-
-By the principle of mathematical induction, the Interaction Net remains semantically equivalent to the original
-Neural Network at every step of the reduction process.
-
-Since Interaction Nets are confluent, the reduced mathematical expression is unique regardless
-of order in which rules are applied.
diff --git a/examples/double_integrator/double_integrator_pretrain_inv.onnx b/examples/double_integrator/double_integrator_pretrain_inv.onnx
new file mode 100644
index 0000000..5fd201c
--- /dev/null
+++ b/examples/double_integrator/double_integrator_pretrain_inv.onnx
Binary files differ
diff --git a/examples/pendulum/pendulum_pretrain_con.onnx b/examples/pendulum/pendulum_pretrain_con.onnx
new file mode 100644
index 0000000..1157a33
--- /dev/null
+++ b/examples/pendulum/pendulum_pretrain_con.onnx
Binary files differ
diff --git a/vein.py b/vein.py
index ca633a2..c42b1a2 100644
--- a/vein.py
+++ b/vein.py
@@ -13,7 +13,6 @@
# If not, see <https://www.gnu.org/licenses/>.
import z3
-import re
import numpy as np
import subprocess
import onnx
@@ -23,6 +22,9 @@ from typing import List, Dict, Optional
import os
import tempfile
import hashlib
+from itertools import count
+import re
+import ast
sat = z3.sat
unsat = z3.unsat
@@ -65,7 +67,7 @@ rules = """
| (q == 1) && (r == 0) => out ~ x
| (q == 1) && (r != 0) => out ~ TermAdd(x, TermConcrete(r))
| (q != 0) && (r == 0) => out ~ TermMul(TermConcrete(q), x)
- | _ => out ~ TermAdd(TermMul(Concrete(q), x), TermConcrete(r));
+ | _ => out ~ TermAdd(TermMul(TermConcrete(q), x), TermConcrete(r));
Concrete(float k) >< Materialize(out) => out ~ (*L)TermConcrete(k);
"""
@@ -74,14 +76,6 @@ _CACHE = {}
def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]] = None) -> str:
# TODO: Add Range agent
_ = bounds
- class NameGen:
- def __init__(self, prefix="v"):
- self.counter = 0
- self.prefix = prefix
- def next(self) -> str:
- name = f"{self.prefix}{self.counter}"
- self.counter += 1
- return name
def get_initializers(graph) -> Dict[str, np.ndarray]:
initializers = {}
@@ -92,81 +86,77 @@ def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]
def get_attrs(node) -> Dict:
return {attr.name: onnx.helper.get_attribute_value(attr) for attr in node.attribute}
- def get_dim(name):
- for i in list(graph.input) + list(graph.output) + list(graph.value_info):
- if i.name == name: return i.type.tensor_type.shape.dim[-1].dim_value
- return None
+ graph, initializers = model.graph, get_initializers(model.graph)
+ counter = count()
+ wire_gen = lambda: f"w{next(counter)}"
+ interactions: Dict[str, List[List[str]]] = {}
+ dims = {i.name: i.type.tensor_type.shape.dim[-1].dim_value for i in list(graph.input) + list(graph.output) + list(graph.value_info)}
+ script = []
def balanced_fanout(agent_name: str, terms: List[str]) -> str:
if not terms: return "Eraser"
if len(terms) == 1: return terms[0]
- nodes = terms
- while len(nodes) > 1:
+ while len(terms) > 1:
next_level = []
- for i in range(0, len(nodes), 2):
- if i + 1 < len(nodes):
- in_w = wire_gen.next()
- script.append(f"{in_w} ~ {agent_name}({nodes[i]}, {nodes[i+1]});")
+ for i in range(0, len(terms), 2):
+ if i + 1 < len(terms):
+ in_w = wire_gen()
+ script.append(f"{in_w} ~ {agent_name}({terms[i]}, {terms[i+1]});")
next_level.append(in_w)
else:
- next_level.append(nodes[i])
- nodes = next_level
- return nodes[0]
+ next_level.append(terms[i])
+ terms = next_level
+ return terms[0]
def balanced_fanin(agent_name: str, terms: List[str]) -> str:
if not terms: return "Eraser"
if len(terms) == 1: return terms[0]
- nodes = terms
- while len(nodes) > 1:
+ while len(terms) > 1:
next_level = []
- for i in range(0, len(nodes), 2):
- if i + 1 < len(nodes):
- res_w = wire_gen.next()
- script.append(f"{nodes[i]} ~ {agent_name}({res_w}, {nodes[i+1]});")
+ for i in range(0, len(terms), 2):
+ if i + 1 < len(terms):
+ res_w = wire_gen()
+ script.append(f"{terms[i]} ~ {agent_name}({res_w}, {terms[i+1]});")
next_level.append(res_w)
else:
- next_level.append(nodes[i])
- nodes = next_level
- return nodes[0]
-
- def op_gemm(node, override_attrs=None):
- attrs = override_attrs if override_attrs is not None else get_attrs(node)
-
- W = initializers[node.input[1]]
- if not attrs.get("transB", 0): W = W.T
- out_dim, in_dim = W.shape
-
- B = initializers[node.input[2]] if len(node.input) > 2 else np.zeros(out_dim)
- alpha, beta = attrs.get("alpha", 1.0), attrs.get("beta", 1.0)
-
- if node.input[0] not in interactions:
- interactions[node.input[0]] = [[] for _ in range(in_dim)]
-
- out_terms = interactions.get(node.output[0]) or [[f"Materialize(result{j})"] for j in range(out_dim)]
-
+ next_level.append(terms[i])
+ terms = next_level
+ return terms[0]
+
+ def gemm(Y, A, B, C, alpha, beta, _, transB):
+ weights = initializers[B]
+ if transB == 0:
+ weights = weights.T
+ out_dim, in_dim = weights.shape
+ biases = initializers[C] if C is not None else None
+ if A not in interactions:
+ interactions[A] = [[] for _ in range(in_dim)]
+ out_terms = interactions.get(Y) or [[f"Materialize(result{j})"] for j in range(out_dim)]
for j in range(out_dim):
sink = balanced_fanout("Dup", out_terms[j])
neuron_terms = []
for i in range(in_dim):
- weight = float(alpha * W[j, i])
+ weight = float(alpha * weights[j, i])
if weight != 0:
- v = wire_gen.next()
- interactions[node.input[0]][i].append(f"Mul({v}, Concrete({weight}))")
+ v = wire_gen()
+ interactions[A][i].append(f"Mul({v}, Concrete({weight}))")
neuron_terms.append(v)
-
- bias_val = float(beta * B[j])
- if bias_val != 0 or not neuron_terms:
- neuron_terms.append(f"Concrete({bias_val})")
-
+ bias = float(beta * biases[j]) if biases is not None else 0.0
+ if bias != 0 or len(neuron_terms) == 0:
+ neuron_terms.append(f"Concrete({bias})")
root = balanced_fanin("Add", neuron_terms)
script.append(f"{root} ~ {sink};")
+ def op_gemm(node):
+ attrs = get_attrs(node)
+ gemm(node.output[0], node.input[0], node.input[1], node.input[2], attrs.get("alpha", 1.0), attrs.get("beta", 1.0), attrs.get("transA", 0), attrs.get("transB", 0))
+
def op_matmul(node):
- op_gemm(node, override_attrs={"alpha": 1.0, "beta": 0.0, "transB": 0})
+ gemm(node.output[0], node.input[0], node.input[1], None, 1.0, 0.0, 0, 0)
def op_relu(node):
- out_name, in_name = node.output[0], node.input[0]
- dim = get_dim(out_name) or 1
+ in_name, out_name = node.input[0], node.output[0]
+ dim = dims.get(out_name) or 1
if in_name not in interactions:
interactions[in_name] = [[] for _ in range(dim)]
@@ -175,24 +165,20 @@ def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]
for i in range(dim):
sink = balanced_fanout("Dup", out_terms[i])
- v = wire_gen.next()
+ v = wire_gen()
interactions[in_name][i].append(f"ReLU({v})")
script.append(f"{v} ~ {sink};")
- def op_flatten(node):
- op_identity(node)
-
- def op_reshape(node):
- op_identity(node)
-
def op_add(node):
- out_name = node.output[0]
in_a, in_b = node.input[0], node.input[1]
+ out_name = node.output[0]
- dim = get_dim(out_name) or get_dim(in_a) or get_dim(in_b) or 1
+ dim = dims.get(out_name) or dims.get(in_a) or dims.get(in_b) or 1
- if in_a not in interactions: interactions[in_a] = [[] for _ in range(dim)]
- if in_b not in interactions: interactions[in_b] = [[] for _ in range(dim)]
+ if in_a not in interactions:
+ interactions[in_a] = [[] for _ in range(dim)]
+ if in_b not in interactions:
+ interactions[in_b] = [[] for _ in range(dim)]
out_terms = interactions.get(out_name) or [[f"Materialize(result{j})"] for j in range(dim)]
@@ -208,24 +194,23 @@ def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]
val = float(a_const.flatten()[i % a_const.size])
interactions[in_b][i].append(f"Add({sink}, Concrete({val}))")
else:
- v_b = wire_gen.next()
+ v_b = wire_gen()
interactions[in_a][i].append(f"Add({sink}, {v_b})")
interactions[in_b][i].append(f"{v_b}")
def op_sub(node):
- out_name = node.output[0]
in_a, in_b = node.input[0], node.input[1]
+ out_name = node.output[0]
- dim = get_dim(out_name) or get_dim(in_a) or get_dim(in_b) or 1
+ dim = dims.get(out_name) or dims.get(in_a) or dims.get(in_b) or 1
- if out_name not in interactions:
- interactions[out_name] = [[f"Materialize(result{i})"] for i in range(dim)]
+ if in_a not in interactions:
+ interactions[in_a] = [[] for _ in range(dim)]
+ if in_b not in interactions:
+ interactions[in_b] = [[] for _ in range(dim)]
out_terms = interactions.get(out_name) or [[f"Materialize(result{j})"] for j in range(dim)]
- if in_a not in interactions: interactions[in_a] = [[] for _ in range(dim)]
- if in_b not in interactions: interactions[in_b] = [[] for _ in range(dim)]
-
b_const = initializers.get(in_b)
a_const = initializers.get(in_a)
@@ -236,47 +221,29 @@ def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]
interactions[in_a][i].append(f"Add({sink}, Concrete({-val}))")
elif a_const is not None:
val = float(a_const.flatten()[i % a_const.size])
- v_b = wire_gen.next()
interactions[in_b][i].append(f"Mul(Add({sink}, Concrete({val})), Concrete(-1.0))")
else:
- v_b = wire_gen.next()
+ v_b = wire_gen()
interactions[in_a][i].append(f"Add({sink}, {v_b})")
interactions[in_b][i].append(f"Mul({v_b}, Concrete(-1.0))")
- def op_slice(node):
- in_name, out_name = node.input[0], node.output[0]
- if out_name in interactions:
- starts = initializers.get(node.input[1])
- steps = initializers.get(node.input[4]) if len(node.input) > 4 else None
-
- start = int(starts.flatten()[0]) if starts is not None else 0
- step = int(steps.flatten()[0]) if steps is not None else 1
-
- in_dim = get_dim(in_name) or 1
- if in_name not in interactions:
- interactions[in_name] = [[] for _ in range(in_dim)]
-
- for i, terms in enumerate(interactions[out_name]):
- input_index = start + (i * step)
- if input_index < in_dim:
- interactions[in_name][input_index].extend(terms)
-
def op_squeeze(node):
op_identity(node)
def op_unsqueeze(node):
op_identity(node)
+ def op_flatten(node):
+ op_identity(node)
+
+ def op_reshape(node):
+ op_identity(node)
+
def op_identity(node):
in_name, out_name = node.input[0], node.output[0]
if out_name in interactions:
interactions[in_name] = interactions[out_name]
-
- graph, initializers = model.graph, get_initializers(model.graph)
- wire_gen = NameGen("w")
- interactions: Dict[str, List[List[str]]] = {}
- script = []
ops = {
"Gemm": op_gemm,
"Relu": op_relu,
@@ -285,7 +252,6 @@ def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]
"MatMul": op_matmul,
"Add": op_add,
"Sub": op_sub,
- "Slice": op_slice,
"Squeeze": op_squeeze,
"Unsqueeze": op_unsqueeze,
"Identity": op_identity
@@ -293,7 +259,7 @@ def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]
if graph.output:
out = graph.output[0].name
- dim = get_dim(out)
+ dim = dims.get(out)
if dim:
interactions[out] = [[f"Materialize(result{i})"] for i in range(dim)]
@@ -304,10 +270,11 @@ def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]
raise RuntimeError(f"Unsupported ONNX operator: {node.op_type}")
if graph.input:
- input = "input" if "input" in interactions else graph.input[0].name
- for i, terms in enumerate(interactions[input]):
- sink = balanced_fanout("Dup", terms)
- script.append(f"{sink} ~ Linear(TermSymbolic(X_{i}), 1.0, 0.0);")
+ for input in graph.input:
+ if input.name in interactions:
+ for i, terms in enumerate(interactions[input.name]):
+ sink = balanced_fanout("Dup", terms)
+ script.append(f"{sink} ~ Linear(TermSymbolic(X_{i}), 1.0, 0.0);")
result_lines = [f"result{i};" for i in range(len(interactions.get(graph.output[0].name, [])))]
return "\n".join(script + result_lines)
@@ -345,60 +312,35 @@ def z3_evaluate(model: str, X: dict):
'TermReLU': TermReLU
}
- def tokenize(s):
- i = 0
- n = len(s)
- while i < n:
- c = s[i]
- if c in '(),':
- yield c
- i += 1
- elif c.isspace():
- i += 1
- else:
- start = i
- while i < n and s[i] not in '(), ' and not s[i].isspace():
- i += 1
- yield s[start:i]
-
- def iterative_eval(tokens_gen):
- stack = [[]]
- for token in tokens_gen:
- if token == '(':
- stack.append([])
- elif token == ')':
- args = stack.pop()
- func_name = stack[-1].pop()
- func = context.get(func_name)
- if not func: raise ValueError(f"Unknown: {func_name}")
- stack[-1].append(func(*args))
- elif token == ',':
- continue
- else:
- if token in context:
- stack[-1].append(token)
- else:
- try:
- stack[-1].append(float(token))
- except ValueError:
- stack[-1].append(token)
- return stack[0][0]
-
exprs = []
+ allowed_calls = set(context.keys())
+ allowed_nodes = (ast.Expression, ast.Call, ast.Name, ast.Load, ast.Constant, ast.UnaryOp, ast.USub)
+ model = re.sub(r'X_\d+', lambda m: f'"{m.group(0)}"', model)
for line in model.splitlines():
- line = line.strip()
- exprs.append(iterative_eval(tokenize(line)))
+ line = line.strip().rstrip(';')
+ if not line:
+ continue
+ tree = ast.parse(line, mode="eval")
+ for node in ast.walk(tree):
+ if not isinstance(node, allowed_nodes):
+ raise ValueError(f"Disallowed syntax: {type(node).__name__}")
+ if isinstance(node, ast.Call):
+ if not isinstance(node.func, ast.Name) or node.func.id not in allowed_calls:
+ raise ValueError("Disallowed function call")
+ if isinstance(node, ast.Constant) and not isinstance(node.value, (int, float, str)):
+ raise ValueError(f"Only numeric constants and string names allowed")
+ exprs.append(eval(compile(tree, "<model>", "eval"), {"__builtins__": {}}, context))
return exprs
-def net(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]] = None):
+def net(model: onnx.ModelProto, X, bounds: Optional[Dict[str, List[float]]] = None):
model_hash = hashlib.sha256(model.SerializeToString()).hexdigest()
- bounds_key = tuple(sorted((k, tuple(v)) for k, v in bounds.items())) if bounds else None
- cache_key = (model_hash, bounds_key)
+ # bounds_key = tuple(sorted((k, tuple(v)) for k, v in bounds.items())) if bounds else None
+ # cache_key = (model_hash, bounds_key)
+ cache_key = model_hash
if cache_key not in _CACHE:
exported = inpla_export(model, bounds)
reduced = inpla_run(exported)
- X = {}
evaluated = z3_evaluate(reduced, X)
_CACHE[cache_key] = evaluated
@@ -408,19 +350,20 @@ def net(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]] = None)
class Solver(z3.Solver):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
- self.bounds: Dict[str, List[float]] = {}
+ # self.bounds: Dict[str, List[float]] = {}
self.pending_nets: List[onnx.ModelProto] = []
+ self.X = {}
def load_smtlib(self, file_path: str):
with open(file_path, "r") as f:
content = f.read()
- for match in re.finditer(r"\(assert\s+\((>=|<=)\s+(X_\d+)\s+([-+]?\d*\.?\d+(?:[eE][-+]?\d+)?)\)\)", content):
- op, var, val = match.groups()
- val = float(val)
- if var not in self.bounds: self.bounds[var] = [float('-inf'), float('inf')]
- if op == ">=": self.bounds[var][0] = val
- else: self.bounds[var][1] = val
+ # for match in re.finditer(r"\(assert\s+\((>=|<=)\s+(X_\d+)\s+([-+]?\d*\.?\d+(?:[eE][-+]?\d+)?)\)\)", content):
+ # op, var, val = match.groups()
+ # val = float(val)
+ # if var not in self.bounds: self.bounds[var] = [float('-inf'), float('inf')]
+ # if op == ">=": self.bounds[var][0] = val
+ # else: self.bounds[var][1] = val
assertions = z3.parse_smt2_string(content)
self.add(assertions)
@@ -433,9 +376,10 @@ class Solver(z3.Solver):
def _process_nets(self):
y_count = 0
for model in self.pending_nets:
- z3_outputs = net(model, bounds=self.bounds)
+ z3_outputs = net(model, self.X)
+
if z3_outputs:
- for _, out_expr in enumerate(z3_outputs):
+ for out_expr in z3_outputs:
y_var = z3.Real(f"Y_{y_count}")
self.add(y_var == out_expr)
y_count += 1
diff --git a/verify_example.py b/verify_example.py
index a6e8839..92641c8 100644
--- a/verify_example.py
+++ b/verify_example.py
@@ -13,22 +13,19 @@ def check_property(onnx_a, onnx_b, smtlib):
result = solver.check()
if result == vein.unsat:
- print("VERIFIED (UNSAT): The networks are equivalent under this property.", end="\n\n")
+ print("VERIFIED (UNSAT): The networks are equivalent under this property.\n")
elif result == vein.sat:
- print("FAILED (SAT): The networks are NOT equivalent.")
- print("Counter-example input:")
- print(solver.model(), end="\n\n")
+ print(f"FAILED (SAT): The networks are NOT equivalent.\nCounter-example input:\n{solver.model()}\n")
# m = solver.model()
# sorted_symbols = sorted([s for s in m.decls() if s.name().startswith("X_")], key=lambda s: s.name())
# for s in sorted_symbols:
# print(f" {s.name()} = {m[s]}")
else:
- print("UNKNOWN", end="\n\n")
+ print("UNKNOWN\n")
if __name__ == "__main__":
if len(sys.argv) <= 1:
- print("Net not provided")
- print("Available Nets: 'xor', 'mnist', 'iris', 'acasxu', 'pendulum', 'double_integrator'")
+ print("Net not provided\nAvailable Nets: 'xor', 'mnist', 'iris', 'acasxu', 'pendulum', 'double_integrator'")
sys.exit()
match sys.argv[1]:
@@ -57,13 +54,13 @@ if __name__ == "__main__":
epsilon = "./examples/ACASXU/ACASXU_epsilon.smtlib"
argmax = "./examples/ACASXU/ACASXU_argmax.smtlib"
case "pendulum":
- net_a = "./examples/pendulum/pendulum_finetune_con.onnx"
+ net_a = "./examples/pendulum/pendulum_pretrain_con.onnx"
net_b = "./examples/pendulum/pendulum_finetune_con.onnx"
strict = "./examples/pendulum/pendulum_strict.smtlib"
epsilon = "./examples/pendulum/pendulum_epsilon.smtlib"
argmax = "./examples/pendulum/pendulum_argmax.smtlib"
case "double_integrator":
- net_a = "./examples/double_integrator/double_integrator_finetune_inv.onnx"
+ net_a = "./examples/double_integrator/double_integrator_pretrain_inv.onnx"
net_b = "./examples/double_integrator/double_integrator_finetune_inv.onnx"
strict = "./examples/double_integrator/double_integrator_strict.smtlib"
epsilon = "./examples/double_integrator/double_integrator_epsilon.smtlib"
@@ -72,7 +69,7 @@ if __name__ == "__main__":
print("Available Nets: 'xor', 'mnist', 'iris', 'acasxu', 'pendulum', 'double_integrator'")
sys.exit()
- print(f"=== Comparing {net_a} and {net_b} ===", end="\n\n")
+ print(f"=== Comparing {net_a} and {net_b} ===\n")
check_property(net_a, net_b, strict)
check_property(net_a, net_b, epsilon)