diff options
| author | ericmarin <maarin.eric@gmail.com> | 2026-06-05 16:40:44 +0200 |
|---|---|---|
| committer | ericmarin <maarin.eric@gmail.com> | 2026-06-10 14:28:22 +0200 |
| commit | d85de781b4a825fe0d1217e1e8ae5ec81dd70202 (patch) | |
| tree | bf7d0539ff7410ca12d53be7a3e5a6530d277530 | |
| parent | fcbbc960f43137aa170b78ba0be2d89aec3bc766 (diff) | |
| download | vein-d85de781b4a825fe0d1217e1e8ae5ec81dd70202.tar.gz vein-d85de781b4a825fe0d1217e1e8ae5ec81dd70202.zip | |
balanced fanout
| -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.py | 81 | ||||
| -rw-r--r-- | verify_example.py | 42 |
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 @@ -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() |
