Source code for boundlab.diff.onnx

"""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"]