aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--examples/ACASXU/ACASXU_argmax.smtlib (renamed from examples/ACASXU/ACASXU_argmax.vnnlib)0
-rw-r--r--examples/ACASXU/ACASXU_epsilon.smtlib (renamed from examples/ACASXU/ACASXU_epsilon.vnnlib)0
-rw-r--r--examples/ACASXU/ACASXU_strict.smtlib (renamed from examples/ACASXU/ACASXU_strict.vnnlib)0
-rw-r--r--examples/double_integrator/double_integrator_argmax.smtlib (renamed from examples/double_integrator/double_integrator_argmax.vnnlib)0
-rw-r--r--examples/double_integrator/double_integrator_epsilon.smtlib (renamed from examples/double_integrator/double_integrator_epsilon.vnnlib)0
-rw-r--r--examples/double_integrator/double_integrator_strict.smtlib (renamed from examples/double_integrator/double_integrator_strict.vnnlib)0
-rw-r--r--examples/iris/iris_argmax.smtlib (renamed from examples/iris/iris_argmax.vnnlib)0
-rw-r--r--examples/iris/iris_epsilon.smtlib (renamed from examples/iris/iris_epsilon.vnnlib)0
-rw-r--r--examples/iris/iris_strict.smtlib (renamed from examples/iris/iris_strict.vnnlib)0
-rw-r--r--examples/mnist/mnist_argmax.smtlib (renamed from examples/mnist/mnist_argmax.vnnlib)0
-rw-r--r--examples/mnist/mnist_epsilon.smtlib (renamed from examples/mnist/mnist_epsilon.vnnlib)0
-rw-r--r--examples/mnist/mnist_strict.smtlib (renamed from examples/mnist/mnist_strict.vnnlib)0
-rw-r--r--examples/pendulum/pendulum_argmax.smtlib (renamed from examples/pendulum/pendulum_argmax.vnnlib)0
-rw-r--r--examples/pendulum/pendulum_epsilon.smtlib (renamed from examples/pendulum/pendulum_epsilon.vnnlib)0
-rw-r--r--examples/pendulum/pendulum_strict.smtlib (renamed from examples/pendulum/pendulum_strict.vnnlib)0
-rw-r--r--examples/xor/xor_argmax.smtlib (renamed from examples/xor/xor_argmax.vnnlib)0
-rw-r--r--examples/xor/xor_epsilon.smtlib (renamed from examples/xor/xor_epsilon.vnnlib)0
-rw-r--r--examples/xor/xor_strict.smtlib (renamed from examples/xor/xor_strict.vnnlib)0
-rw-r--r--vein.py81
-rw-r--r--verify_example.py42
20 files changed, 63 insertions, 60 deletions
diff --git a/examples/ACASXU/ACASXU_argmax.vnnlib b/examples/ACASXU/ACASXU_argmax.smtlib
index 3009eef..3009eef 100644
--- a/examples/ACASXU/ACASXU_argmax.vnnlib
+++ b/examples/ACASXU/ACASXU_argmax.smtlib
diff --git a/examples/ACASXU/ACASXU_epsilon.vnnlib b/examples/ACASXU/ACASXU_epsilon.smtlib
index 2cbcd36..2cbcd36 100644
--- a/examples/ACASXU/ACASXU_epsilon.vnnlib
+++ b/examples/ACASXU/ACASXU_epsilon.smtlib
diff --git a/examples/ACASXU/ACASXU_strict.vnnlib b/examples/ACASXU/ACASXU_strict.smtlib
index 1e5b8e2..1e5b8e2 100644
--- a/examples/ACASXU/ACASXU_strict.vnnlib
+++ b/examples/ACASXU/ACASXU_strict.smtlib
diff --git a/examples/double_integrator/double_integrator_argmax.vnnlib b/examples/double_integrator/double_integrator_argmax.smtlib
index f0a72e9..f0a72e9 100644
--- a/examples/double_integrator/double_integrator_argmax.vnnlib
+++ b/examples/double_integrator/double_integrator_argmax.smtlib
diff --git a/examples/double_integrator/double_integrator_epsilon.vnnlib b/examples/double_integrator/double_integrator_epsilon.smtlib
index f5c4ee6..f5c4ee6 100644
--- a/examples/double_integrator/double_integrator_epsilon.vnnlib
+++ b/examples/double_integrator/double_integrator_epsilon.smtlib
diff --git a/examples/double_integrator/double_integrator_strict.vnnlib b/examples/double_integrator/double_integrator_strict.smtlib
index d3c8c3e..d3c8c3e 100644
--- a/examples/double_integrator/double_integrator_strict.vnnlib
+++ b/examples/double_integrator/double_integrator_strict.smtlib
diff --git a/examples/iris/iris_argmax.vnnlib b/examples/iris/iris_argmax.smtlib
index ec72109..ec72109 100644
--- a/examples/iris/iris_argmax.vnnlib
+++ b/examples/iris/iris_argmax.smtlib
diff --git a/examples/iris/iris_epsilon.vnnlib b/examples/iris/iris_epsilon.smtlib
index df691c4..df691c4 100644
--- a/examples/iris/iris_epsilon.vnnlib
+++ b/examples/iris/iris_epsilon.smtlib
diff --git a/examples/iris/iris_strict.vnnlib b/examples/iris/iris_strict.smtlib
index 78d01fe..78d01fe 100644
--- a/examples/iris/iris_strict.vnnlib
+++ b/examples/iris/iris_strict.smtlib
diff --git a/examples/mnist/mnist_argmax.vnnlib b/examples/mnist/mnist_argmax.smtlib
index 4c7f0c9..4c7f0c9 100644
--- a/examples/mnist/mnist_argmax.vnnlib
+++ b/examples/mnist/mnist_argmax.smtlib
diff --git a/examples/mnist/mnist_epsilon.vnnlib b/examples/mnist/mnist_epsilon.smtlib
index ea76779..ea76779 100644
--- a/examples/mnist/mnist_epsilon.vnnlib
+++ b/examples/mnist/mnist_epsilon.smtlib
diff --git a/examples/mnist/mnist_strict.vnnlib b/examples/mnist/mnist_strict.smtlib
index 356f176..356f176 100644
--- a/examples/mnist/mnist_strict.vnnlib
+++ b/examples/mnist/mnist_strict.smtlib
diff --git a/examples/pendulum/pendulum_argmax.vnnlib b/examples/pendulum/pendulum_argmax.smtlib
index c11dc0b..c11dc0b 100644
--- a/examples/pendulum/pendulum_argmax.vnnlib
+++ b/examples/pendulum/pendulum_argmax.smtlib
diff --git a/examples/pendulum/pendulum_epsilon.vnnlib b/examples/pendulum/pendulum_epsilon.smtlib
index 8209db5..8209db5 100644
--- a/examples/pendulum/pendulum_epsilon.vnnlib
+++ b/examples/pendulum/pendulum_epsilon.smtlib
diff --git a/examples/pendulum/pendulum_strict.vnnlib b/examples/pendulum/pendulum_strict.smtlib
index d1c1167..d1c1167 100644
--- a/examples/pendulum/pendulum_strict.vnnlib
+++ b/examples/pendulum/pendulum_strict.smtlib
diff --git a/examples/xor/xor_argmax.vnnlib b/examples/xor/xor_argmax.smtlib
index d21119b..d21119b 100644
--- a/examples/xor/xor_argmax.vnnlib
+++ b/examples/xor/xor_argmax.smtlib
diff --git a/examples/xor/xor_epsilon.vnnlib b/examples/xor/xor_epsilon.smtlib
index 427243e..427243e 100644
--- a/examples/xor/xor_epsilon.vnnlib
+++ b/examples/xor/xor_epsilon.smtlib
diff --git a/examples/xor/xor_strict.vnnlib b/examples/xor/xor_strict.smtlib
index bead476..bead476 100644
--- a/examples/xor/xor_strict.vnnlib
+++ b/examples/xor/xor_strict.smtlib
diff --git a/vein.py b/vein.py
index baa9da5..ca633a2 100644
--- a/vein.py
+++ b/vein.py
@@ -61,12 +61,12 @@ rules = """
| k > 0 => out ~ (*L)Concrete(k)
| _ => out ~ Concrete(0);
Linear(x, float q, float r) >< Materialize(out)
- | (q == 0) => out ~ Concrete(r), x ~ Eraser
+ | (q == 0) => out ~ TermConcrete(r), x ~ Eraser
| (q == 1) && (r == 0) => out ~ x
- | (q == 1) && (r != 0) => out ~ TermAdd(x, Concrete(r))
- | (q != 0) && (r == 0) => out ~ TermMul(Concrete(q), x)
- | _ => out ~ TermAdd(TermMul(Concrete(q), x), Concrete(r));
- Concrete(float k) >< Materialize(out) => out ~ (*L)Concrete(k);
+ | (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));
+ Concrete(float k) >< Materialize(out) => out ~ (*L)TermConcrete(k);
"""
_CACHE = {}
@@ -97,36 +97,37 @@ def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]
if i.name == name: return i.type.tensor_type.shape.dim[-1].dim_value
return None
- def flatten_nest(agent_name: str, terms: List[str]) -> str:
+ def balanced_fanout(agent_name: str, terms: List[str]) -> str:
if not terms: return "Eraser"
if len(terms) == 1: return terms[0]
- current = terms[0]
- for i in range(1, len(terms)):
- wire = wire_gen.next()
- script.append(f"{wire} ~ {agent_name}({current}, {terms[i]});")
- current = wire
- return current
-
- def balance_add(terms: List[str], sink: str):
- if not terms:
- script.append(f"{sink} ~ Eraser;")
- return
- if len(terms) == 1:
- script.append(f"{sink} ~ {terms[0]};")
- return
+ nodes = terms
+ while len(nodes) > 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]});")
+ next_level.append(in_w)
+ else:
+ next_level.append(nodes[i])
+ nodes = next_level
+ return nodes[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:
next_level = []
for i in range(0, len(nodes), 2):
if i + 1 < len(nodes):
- wire_out = wire_gen.next()
- script.append(f"{nodes[i]} ~ Add({wire_out}, {nodes[i+1]});")
- next_level.append(wire_out)
+ res_w = wire_gen.next()
+ script.append(f"{nodes[i]} ~ {agent_name}({res_w}, {nodes[i+1]});")
+ next_level.append(res_w)
else:
next_level.append(nodes[i])
nodes = next_level
- script.append(f"{nodes[0]} ~ {sink};")
+ return nodes[0]
def op_gemm(node, override_attrs=None):
attrs = override_attrs if override_attrs is not None else get_attrs(node)
@@ -144,7 +145,7 @@ def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]
out_terms = interactions.get(node.output[0]) or [[f"Materialize(result{j})"] for j in range(out_dim)]
for j in range(out_dim):
- sink = flatten_nest("Dup", out_terms[j])
+ sink = balanced_fanout("Dup", out_terms[j])
neuron_terms = []
for i in range(in_dim):
weight = float(alpha * W[j, i])
@@ -157,7 +158,8 @@ def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]
if bias_val != 0 or not neuron_terms:
neuron_terms.append(f"Concrete({bias_val})")
- balance_add(neuron_terms, sink)
+ root = balanced_fanin("Add", neuron_terms)
+ script.append(f"{root} ~ {sink};")
def op_matmul(node):
op_gemm(node, override_attrs={"alpha": 1.0, "beta": 0.0, "transB": 0})
@@ -172,7 +174,7 @@ def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]
out_terms = interactions.get(node.output[0]) or [[f"Materialize(result{j})"] for j in range(dim)]
for i in range(dim):
- sink = flatten_nest("Dup", out_terms[i])
+ sink = balanced_fanout("Dup", out_terms[i])
v = wire_gen.next()
interactions[in_name][i].append(f"ReLU({v})")
script.append(f"{v} ~ {sink};")
@@ -198,7 +200,7 @@ def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]
a_const = initializers.get(in_a)
for i in range(dim):
- sink = flatten_nest("Dup", out_terms[i])
+ sink = balanced_fanout("Dup", out_terms[i])
if b_const is not None:
val = float(b_const.flatten()[i % b_const.size])
interactions[in_a][i].append(f"Add({sink}, Concrete({val}))")
@@ -228,17 +230,18 @@ def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]
a_const = initializers.get(in_a)
for i in range(dim):
- sink = flatten_nest("Dup", out_terms[i])
+ sink = balanced_fanout("Dup", out_terms[i])
if b_const is not None:
val = float(b_const.flatten()[i % b_const.size])
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])
- interactions[in_b][i].append(f"Add({sink}, Concrete({-val}))")
+ 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()
- interactions[in_a][i].append(f"Add({sink}, Mul({v_b}, Concrete(-1.0)))")
- interactions[in_b][i].append(f"{v_b}")
+ 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]
@@ -303,8 +306,8 @@ def inpla_export(model: onnx.ModelProto, bounds: Optional[Dict[str, List[float]]
if graph.input:
input = "input" if "input" in interactions else graph.input[0].name
for i, terms in enumerate(interactions[input]):
- sink = flatten_nest("Dup", terms)
- script.append(f"{sink} ~ Linear(Symbolic(X_{i}), 1.0, 0.0);")
+ 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)
@@ -325,18 +328,18 @@ def inpla_run(model: str) -> str:
def z3_evaluate(model: str, X: dict):
- def Symbolic(id):
+ def TermSymbolic(id):
if id not in X:
X[id] = z3.Real(id)
return X[id]
- def Concrete(val): return z3.RealVal(val)
+ def TermConcrete(val): return z3.RealVal(val)
def TermAdd(a, b): return a + b
def TermMul(a, b): return a * b
def TermReLU(x): return z3.If(x > 0, x, 0)
context = {
- 'Concrete': Concrete,
- 'Symbolic': Symbolic,
+ 'TermConcrete': TermConcrete,
+ 'TermSymbolic': TermSymbolic,
'TermAdd': TermAdd,
'TermMul': TermMul,
'TermReLU': TermReLU
@@ -408,7 +411,7 @@ class Solver(z3.Solver):
self.bounds: Dict[str, List[float]] = {}
self.pending_nets: List[onnx.ModelProto] = []
- def load_vnnlib(self, file_path: str):
+ def load_smtlib(self, file_path: str):
with open(file_path, "r") as f:
content = f.read()
diff --git a/verify_example.py b/verify_example.py
index 1bf00fa..a6e8839 100644
--- a/verify_example.py
+++ b/verify_example.py
@@ -1,14 +1,14 @@
import sys
import vein
-def check_property(onnx_a, onnx_b, vnnlib):
+def check_property(onnx_a, onnx_b, smtlib):
solver = vein.Solver()
- print(f"--- Checking {vnnlib} ---")
+ print(f"--- Checking {smtlib} ---")
solver.load_onnx(onnx_a)
solver.load_onnx(onnx_b)
- solver.load_vnnlib(vnnlib)
+ solver.load_smtlib(smtlib)
result = solver.check()
@@ -35,39 +35,39 @@ if __name__ == "__main__":
case "xor":
net_a = "./examples/xor/xor_a.onnx"
net_b = "./examples/xor/xor_b.onnx"
- strict = "./examples/xor/xor_strict.vnnlib"
- epsilon = "./examples/xor/xor_epsilon.vnnlib"
- argmax = "./examples/xor/xor_argmax.vnnlib"
+ strict = "./examples/xor/xor_strict.smtlib"
+ epsilon = "./examples/xor/xor_epsilon.smtlib"
+ argmax = "./examples/xor/xor_argmax.smtlib"
case "mnist":
net_a = "./examples/mnist/mnist_a.onnx"
net_b = "./examples/mnist/mnist_b.onnx"
- strict = "./examples/mnist/mnist_strict.vnnlib"
- epsilon = "./examples/mnist/mnist_epsilon.vnnlib"
- argmax = "./examples/mnist/mnist_argmax.vnnlib"
+ strict = "./examples/mnist/mnist_strict.smtlib"
+ epsilon = "./examples/mnist/mnist_epsilon.smtlib"
+ argmax = "./examples/mnist/mnist_argmax.smtlib"
case "iris":
net_a = "./examples/iris/iris_a.onnx"
net_b = "./examples/iris/iris_b.onnx"
- strict = "./examples/iris/iris_strict.vnnlib"
- epsilon = "./examples/iris/iris_epsilon.vnnlib"
- argmax = "./examples/iris/iris_argmax.vnnlib"
+ strict = "./examples/iris/iris_strict.smtlib"
+ epsilon = "./examples/iris/iris_epsilon.smtlib"
+ argmax = "./examples/iris/iris_argmax.smtlib"
case "acasxu":
net_a = "./examples/ACASXU/ACASXU_run2a_1_1_batch_2000.onnx"
net_b = "./examples/ACASXU/ACASXU_run2a_1_1_batch_2000.onnx"
- strict = "./examples/ACASXU/ACASXU_strict.vnnlib"
- epsilon = "./examples/ACASXU/ACASXU_epsilon.vnnlib"
- argmax = "./examples/ACASXU/ACASXU_argmax.vnnlib"
+ strict = "./examples/ACASXU/ACASXU_strict.smtlib"
+ epsilon = "./examples/ACASXU/ACASXU_epsilon.smtlib"
+ argmax = "./examples/ACASXU/ACASXU_argmax.smtlib"
case "pendulum":
net_a = "./examples/pendulum/pendulum_finetune_con.onnx"
net_b = "./examples/pendulum/pendulum_finetune_con.onnx"
- strict = "./examples/pendulum/pendulum_strict.vnnlib"
- epsilon = "./examples/pendulum/pendulum_epsilon.vnnlib"
- argmax = "./examples/pendulum/pendulum_argmax.vnnlib"
+ 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_b = "./examples/double_integrator/double_integrator_finetune_inv.onnx"
- strict = "./examples/double_integrator/double_integrator_strict.vnnlib"
- epsilon = "./examples/double_integrator/double_integrator_epsilon.vnnlib"
- argmax = "./examples/double_integrator/double_integrator_argmax.vnnlib"
+ strict = "./examples/double_integrator/double_integrator_strict.smtlib"
+ epsilon = "./examples/double_integrator/double_integrator_epsilon.smtlib"
+ argmax = "./examples/double_integrator/double_integrator_argmax.smtlib"
case _:
print("Available Nets: 'xor', 'mnist', 'iris', 'acasxu', 'pendulum', 'double_integrator'")
sys.exit()