"""Merge two ONNX models into one paired graph for differential interpretation.
:func:`diff_net` walks two structurally identical graphs in lock-step and,
wherever both sides read a parameter from an initializer, replaces that input
by a shared ``boundlab::DiffPair`` value. At concrete-tensor runtime the
merged model still evaluates network 1; under
:data:`boundlab.diff.zono3.interpret` the paired initializers become
:class:`~boundlab.diff.expr.DiffExpr2` values that carry both branches through
the rest of the graph.
"""
from __future__ import annotations
import copy
from itertools import zip_longest
from pathlib import Path
import onnx_ir as ir
from boundlab.diff import ops as _ops # noqa: F401 (registers the ONNX op names)
def _load_onnx(model: ir.Model | str | Path) -> ir.Model:
if isinstance(model, (str, Path)):
return ir.load(str(model))
if isinstance(model, ir.Model):
return model
raise TypeError(f"Expected a path or onnx_ir.Model, got {type(model)}.")
def _value_ref(name: str) -> ir.Value:
return ir.Value(name=name)
def _new_initializer(name: str, source: ir.Value) -> ir.Value:
return ir.Value(
name=name,
shape=copy.deepcopy(source.shape),
type=copy.deepcopy(source.type),
doc_string=source.doc_string,
const_value=copy.deepcopy(source.const_value),
metadata_props=copy.deepcopy(source.metadata_props),
)
[docs]
def diff_net(
net1: ir.Model | str | Path,
net2: ir.Model | str | Path,
) -> ir.Model:
"""Pair the initializers of two structurally identical ONNX models.
Parameters
----------
net1, net2:
Models with the same node sequence and input arities. Each may be an
:class:`onnx_ir.Model`, a ``str``, or a :class:`pathlib.Path`.
Returns
-------
A merged :class:`onnx_ir.Model` whose paired parameters flow through
``boundlab::DiffPair`` nodes.
Raises
------
ValueError
If the two graphs differ in node count, node type, or input arity.
"""
net1 = _load_onnx(net1)
net2 = _load_onnx(net2)
merged = net1.clone()
init1 = merged.graph.initializers
init2 = net2.graph.initializers
nodes1 = list(merged.graph)
nodes2 = list(net2.graph)
if len(nodes1) != len(nodes2):
raise ValueError(
f"Networks have different numbers of nodes: "
f"{len(nodes1)} vs {len(nodes2)}."
)
new_nodes: list[ir.Node] = []
new_initializers: dict[str, ir.Value] = {}
cloned_init2: dict[str, str] = {}
paired_name: dict[tuple[str, str], str] = {}
def paired_input(name1: str, name2: str) -> ir.Value:
key = (name1, name2)
if key in paired_name:
return _value_ref(paired_name[key])
if name2 not in cloned_init2:
clone_name = f"_diff_init2_{len(cloned_init2)}"
new_initializers[clone_name] = _new_initializer(clone_name, init2[name2])
cloned_init2[name2] = clone_name
clone_name = cloned_init2[name2]
out_name = f"_diff_paired_{len(paired_name)}"
paired_name[key] = out_name
new_nodes.append(
ir.Node(
"boundlab",
"DiffPair",
[_value_ref(name1), _value_ref(clone_name)],
outputs=[ir.Value(name=out_name)],
)
)
return _value_ref(out_name)
for index, (node, node2) in enumerate(zip(nodes1, nodes2)):
if node.op_type != node2.op_type or node.domain != node2.domain:
raise ValueError(
f"Node mismatch at index {index}: "
f"({node.domain}::{node.op_type}) vs ({node2.domain}::{node2.op_type})."
)
if len(node.inputs) != len(node2.inputs):
raise ValueError(
f"Node input-arity mismatch at index {index}: "
f"{len(node.inputs)} vs {len(node2.inputs)}."
)
remapped: list[ir.Value | None] = []
for input1, input2 in zip_longest(node.inputs, node2.inputs):
if input1 is None:
remapped.append(None)
elif input2 is None or input2.name is None:
# net2 omits this optional input: keep net1's value unpaired.
remapped.append(_value_ref(input1.name))
elif input1.name in init1 and input2.name in init2:
remapped.append(paired_input(input1.name, input2.name))
else:
remapped.append(_value_ref(input1.name))
new_nodes.append(
ir.Node(
node.domain,
node.op_type,
remapped,
attributes={
name: copy.deepcopy(value)
for name, value in node.attributes.items()
},
outputs=[ir.Value(name=output.name) for output in node.outputs],
overload=node.overload,
version=node.version,
name=node.name,
doc_string=node.doc_string,
metadata_props=copy.deepcopy(node.metadata_props),
)
)
merged.graph.remove(list(merged.graph), safe=False)
merged.graph.extend(new_nodes)
merged.graph.initializers.update(new_initializers)
if paired_name:
for imports in (merged.graph.opset_imports, merged.opset_imports):
imports["boundlab"] = max(imports.get("boundlab", 0), 1)
return merged
__all__ = ["diff_net"]