# SPDX-License-Identifier: Apache-2.0 """Weights of the Diffusion Planner v5.0 export, read as DATA from the three ONNX files and the param JSON. Every tensor is addressed through the graph node that consumes it (``ttaw.weights.OnnxWeights``), never by initializer name: most MatMul weights are anonymous (``onnx::MatMul_4039``), several biases are deduplicated across modules (``static_encoder/projection/fc2`` reads ``...fc1.bias``; ``route_encoder/attribute_emb`` reads ``...speed_limit_emb.bias``), the decoder cross-attention K/V of all three blocks is one unnamed ``[256, 1536]`` MatMul followed by a Split, and the fusion / cross-attention Q-K-V biases and weights are constant-folded tensors (SPEC 6.2). The result is one flat ``{canonical name: float32 array}`` dict that the CPU reference and the ttnn graph both consume: - linear layers: ``.w`` ``[in, out]`` (``y = x @ w + b``) and ``.b`` ``[out]``; - LayerNorms: ``.gamma`` / ``.beta``; - attention: ``...attn.q`` / ``.kv`` (fusion: Q from LN(x), K|V from x), ``...attn.qkv`` (DiT self-attention), ``...cross_attn.q`` / ``.kv`` (K|V of one block, a column block of the fused cross K/V MatMul), ``...out``; - embeddings: ``decoder.dit.agent_embedding`` ``[2, 256]`` (ego, neighbour), ``encoder.route_position_embedding`` ``[25, 256]``, ``encoder._encoder.unknown_speed_emb`` ``[128]``. There is no BatchNorm in this network, so nothing is folded here. The exact rewrites the TT port applies on top of these tensors (per-step adaLN tables folded into the LayerNorm affine, hoisted cross K/V, the pad-relative fp32 pre-projection island) are in :mod:`.rewrites`, computed from this dict in float64 with one final rounding. The loader also checks the invariants the reference and the port rely on (LayerNorm epsilon 1e-5 everywhere, exact GELU in the encoder and the decoder pre-projection / t-embedder, tanh GELU in the DiT MLPs and the final projection, attention scale 1/sqrt(32), the ``-inf`` key mask), so a different export fails loudly instead of silently. """ from __future__ import annotations import json import math from dataclasses import dataclass, field from pathlib import Path from typing import Any, Dict, List, Mapping, Optional, Tuple import numpy as np from . import config as C from ..ttaw.weights import OnnxWeights, file_sha256 __all__ = ["PlannerWeights", "load_weights", "load_param_json", "Normalization", "find_weights_dir", "MIXER_ENCODERS", "SMALL_ENCODERS"] # categories with an MLP-Mixer trunk -> ONNX module prefix MIXER_ENCODERS = {"ego": "ego_encoder", "neighbor": "neighbor_encoder", "lane": "lane_encoder", "route": "route_encoder", "polygon": "polygon_encoder", "line_string": "line_string_encoder"} # categories encoded by a channel MLP + LayerNorm + projection (goal pose, ego shape, turn indicators) SMALL_ENCODERS = {"goal": "goal_pose_encoder", "ego_shape": "ego_shape_encoder", "turn": "turn_indicator_encoder"} # ------------------------------------------------------------------------------------------------ param JSON @dataclass(frozen=True) class Normalization: """``observation_normalizer`` (per input tensor: mean / std over the last dim) and ``state_normalizer`` (per agent and pose dim) of ``diffusion_planner.param.json`` (PKG/include/.../utils/arg_reader.hpp:80-140).""" observation: Dict[str, Tuple[np.ndarray, np.ndarray]] state_mean: np.ndarray # [321, 4] (or [4]) state_std: np.ndarray major_version: int args: Dict[str, Any] = field(default_factory=dict) def state(self) -> Tuple[np.ndarray, np.ndarray]: """``(mean, std)`` broadcastable to ``[321, T, 4]``.""" return _per_agent(self.state_mean), _per_agent(self.state_std) def _per_agent(v: np.ndarray) -> np.ndarray: v = np.asarray(v, np.float32).reshape(-1) if v.size == C.POSE_DIM: return v.reshape(1, 1, C.POSE_DIM) if v.size == C.MAX_NUM_AGENTS * C.POSE_DIM: return v.reshape(C.MAX_NUM_AGENTS, 1, C.POSE_DIM) raise ValueError(f"unsupported state normalizer size {v.size}") def load_param_json(path: Path) -> Normalization: """Parse ``diffusion_planner.param.json`` like ``arg_reader.hpp``; refuses a major version other than 5.""" with open(path) as f: j = json.load(f) major = int(j.get("major_version", -1)) if major != C.WEIGHT_MAJOR_VERSION: raise ValueError(f"{path}: major_version {major}, this port needs {C.WEIGHT_MAJOR_VERSION} (constants.hpp:22)") obs = {} for key, v in j["observation_normalizer"].items(): mean = np.asarray(v.get("mean", []), np.float32).reshape(-1) std = np.asarray(v.get("std", []), np.float32).reshape(-1) if mean.shape != std.shape: raise ValueError(f"{path}: normalizer {key!r} mean / std sizes differ") obs[key] = (mean, std) sn = j["state_normalizer"] args = {k: v for k, v in j.items() if k not in ("observation_normalizer", "state_normalizer")} return Normalization(obs, np.asarray(sn["mean"], np.float32), np.asarray(sn["std"], np.float32), major, args) # ------------------------------------------------------------------------------------------------ ONNX weights @dataclass class PlannerWeights: """The flat canonical parameter dict plus provenance (file sha256) and the export facts checked at load.""" params: Dict[str, np.ndarray] normalization: Normalization sha256: Dict[str, str] facts: Dict[str, Any] path: Path def __getitem__(self, name: str) -> np.ndarray: return self.params[name] def linear(self, name: str) -> Tuple[np.ndarray, np.ndarray]: return self.params[f"{name}.w"], self.params[f"{name}.b"] def num_parameters(self) -> int: return int(sum(v.size for v in self.params.values())) def to_torch(self, dtype: Any = None) -> Dict[str, Any]: """``{name: torch.Tensor}`` (float32 by default; float64 for the fp64 rewrites).""" import torch dt = dtype or torch.float32 return {k: torch.from_numpy(np.ascontiguousarray(v)).to(dt) for k, v in self.params.items()} class _Reader: """Node-addressed access to one ONNX file with the conventions of this export.""" def __init__(self, path: Path): self.w = OnnxWeights(path) self.facts: Dict[str, Any] = {"gelu": {}, "ln_eps": set()} def has_node(self, name: str) -> bool: try: self.w.node(name) return True except KeyError: return False def const_input(self, node_name: str) -> np.ndarray: """The single constant input of a binary node (the bias of a MatMul + Add pair).""" node = self.w.node(node_name) consts = [t for t in node.inputs if t and self.w.has(t)] if len(consts) != 1: raise ValueError(f"{node_name}: expected one constant input, found {len(consts)}") return np.asarray(self.w.array(consts[0]), np.float32) def bias_after(self, matmul_name: str) -> np.ndarray: add = self.w.consumer_of(self.w.node(matmul_name).outputs[0], "Add") return self.const_input(add.name) def linear(self, path: str) -> Tuple[np.ndarray, np.ndarray]: """``(w [in, out], b [out])`` of the torch ``nn.Linear`` exported under ````: either ``/MatMul`` followed by an ``Add`` (3-D inputs) or ``/Gemm`` (2-D inputs); exactly one of the two must exist.""" has_mm, has_gemm = self.has_node(f"{path}/MatMul"), self.has_node(f"{path}/Gemm") if has_mm == has_gemm: raise KeyError(f"{path}: expected exactly one of MatMul / Gemm, found {has_mm=} {has_gemm=}") if has_mm: w = np.asarray(self.w.matmul_weight(f"{path}/MatMul"), np.float32) return w, self.bias_after(f"{path}/MatMul") return self.gemm(f"{path}/Gemm") def gemm(self, node_name: str) -> Tuple[np.ndarray, np.ndarray]: """``(w [in, out], b [out])`` of a ``Gemm`` node (``y = x @ W^T + b`` with ``transB = 1``).""" g = self.w.gemm(node_name) if g.trans_a or g.alpha != 1.0 or g.beta != 1.0 or g.bias is None: raise ValueError(f"{node_name}: unexpected attributes {g.trans_a=} {g.alpha=} {g.beta=}") w = np.asarray(g.weight, np.float32) return (w.T if g.trans_b else w).copy(), np.asarray(g.bias, np.float32).reshape(-1) def layer_norm(self, path: str) -> Tuple[np.ndarray, np.ndarray]: node = self.w.node(f"{path}/LayerNormalization") self.facts["ln_eps"].add(float(node.attrs.get("epsilon", 1e-5))) if int(node.attrs.get("axis", -1)) != -1: raise ValueError(f"{path}: LayerNormalization over axis {node.attrs.get('axis')}") return (np.asarray(self.w.param(node.name, 1), np.float32), np.asarray(self.w.param(node.name, 2), np.float32)) def gelu(self, path: str) -> str: node = self.w.node(f"{path}/Gelu") mode = str(node.attrs.get("approximate", "none")) self.facts["gelu"][path] = mode return mode def scalar(self, node_name: str) -> float: return float(np.asarray(self.const_input(node_name)).reshape(-1)[0]) def _put_linear(params: Dict[str, np.ndarray], name: str, wb: Tuple[np.ndarray, np.ndarray]) -> None: params[f"{name}.w"], params[f"{name}.b"] = wb def _put_ln(params: Dict[str, np.ndarray], name: str, gb: Tuple[np.ndarray, np.ndarray]) -> None: params[f"{name}.gamma"], params[f"{name}.beta"] = gb def _mlp(r: _Reader, params: Dict[str, np.ndarray], onnx_path: str, name: str, gelu: str) -> None: _put_linear(params, f"{name}.fc1", r.linear(f"{onnx_path}/fc1")) _put_linear(params, f"{name}.fc2", r.linear(f"{onnx_path}/fc2")) mode = r.gelu(f"{onnx_path}/act") if mode != gelu: raise ValueError(f"{onnx_path}: GELU approximate={mode!r}, expected {gelu!r}") def _read_encoder(path: Path, params: Dict[str, np.ndarray]) -> Dict[str, Any]: r = _Reader(path) for cat, mod in {**MIXER_ENCODERS, **SMALL_ENCODERS}.items(): P, N = f"/encoder/{mod}", f"encoder.{mod}" _mlp(r, params, f"{P}/channel_pre_project", f"{N}.channel_pre_project", "none") if cat in MIXER_ENCODERS: _mlp(r, params, f"{P}/token_pre_project", f"{N}.token_pre_project", "none") for i in range(C.MIXER_DEPTH): B, BN = f"{P}/blocks.{i}", f"{N}.blocks.{i}" _put_ln(params, f"{BN}.norm1", r.layer_norm(f"{B}/norm1")) _mlp(r, params, f"{B}/tokens_mlp", f"{BN}.tokens_mlp", "none") _put_ln(params, f"{BN}.norm2", r.layer_norm(f"{B}/norm2")) _mlp(r, params, f"{B}/channels_mlp", f"{BN}.channels_mlp", "none") _put_ln(params, f"{N}.norm", r.layer_norm(f"{P}/norm")) _mlp(r, params, f"{P}/emb_project", f"{N}.emb_project", "none") _put_linear(params, "encoder.neighbor_encoder.type_emb", r.linear("/encoder/neighbor_encoder/type_emb")) for mod in ("lane_encoder", "route_encoder"): _put_linear(params, f"encoder.{mod}.speed_limit_emb", r.linear(f"/encoder/{mod}/speed_limit_emb")) _put_linear(params, f"encoder.{mod}.attribute_emb", r.linear(f"/encoder/{mod}/attribute_emb")) unk = r.w.param(f"/encoder/{mod}/unknown_speed_emb/Gather", 0) params[f"encoder.{mod}.unknown_speed_emb"] = np.asarray(unk, np.float32).reshape(-1) _mlp(r, params, "/encoder/static_encoder/projection", "encoder.static_encoder.projection", "none") _put_linear(params, "encoder.pos_emb", r.linear("/encoder/pos_emb")) rpe = np.asarray(r.w.param("/encoder/Slice_5", 0), np.float32) params["encoder.route_position_embedding"] = rpe.reshape(C.NUM_SEGMENTS_IN_ROUTE, C.HIDDEN_DIM) scales, fills = set(), set() for i in range(C.FUSION_DEPTH): B, BN = f"/encoder/fusion/blocks.{i}", f"encoder.fusion.blocks.{i}" _put_ln(params, f"{BN}.norm1", r.layer_norm(f"{B}/norm1")) q_w = np.asarray(r.w.matmul_weight(f"{B}/attn/MatMul"), np.float32) _put_linear(params, f"{BN}.attn.q", (q_w, r.bias_after(f"{B}/attn/MatMul"))) kv_w = np.asarray(r.w.matmul_weight(f"{B}/attn/MatMul_1"), np.float32) _put_linear(params, f"{BN}.attn.kv", (kv_w, r.bias_after(f"{B}/attn/MatMul_1"))) _put_linear(params, f"{BN}.attn.out", r.gemm(f"{B}/attn/Gemm")) _put_ln(params, f"{BN}.norm2", r.layer_norm(f"{B}/norm2")) _mlp(r, params, f"{B}/mlp", f"{BN}.mlp", "none") scales.add(r.scalar(f"{B}/attn/Mul_3")) # the key-padding bias Where(mask, -inf, 0) is built once in block 0 and shared by the six blocks fills.add(float(np.asarray(r.w.param("/encoder/fusion/blocks.0/attn/Where", 1)).reshape(-1)[0])) _put_ln(params, "encoder.fusion.norm", r.layer_norm("/encoder/fusion/norm")) return {"gelu": r.facts["gelu"], "ln_eps": r.facts["ln_eps"], "attn_scale": scales, "mask_fill": fills, "sha256": r.w.sha256, "nodes": len(r.w.nodes())} def _read_decoder(path: Path, params: Dict[str, np.ndarray]) -> Dict[str, Any]: r = _Reader(path) _mlp(r, params, "/dit/preproj", "decoder.dit.preproj", "none") _mlp(r, params, "/dit/t_embedder", "decoder.dit.t_embedder", "none") ego_row = np.asarray(r.w.param("/dit/Concat_3", 0), np.float32).reshape(-1) expand = r.w.producer(r.w.node("/dit/Concat_3").inputs[1]) nb_row = np.asarray(r.w.param(expand.name, 0), np.float32).reshape(-1) params["decoder.dit.agent_embedding"] = np.stack([ego_row, nb_row]) scales, fills = set(), set() for i in range(C.DIT_DEPTH): B, BN = f"/dit/blocks.{i}", f"decoder.dit.blocks.{i}" _put_linear(params, f"{BN}.adaLN_modulation", r.linear(f"{B}/adaLN_modulation/adaLN_modulation.1")) for n in ("norm1", "norm2", "norm3", "norm4"): _put_ln(params, f"{BN}.{n}", r.layer_norm(f"{B}/{n}")) qkv_w = np.asarray(r.w.matmul_weight(f"{B}/attn/MatMul"), np.float32) _put_linear(params, f"{BN}.attn.qkv", (qkv_w, r.bias_after(f"{B}/attn/MatMul"))) _put_linear(params, f"{BN}.attn.out", r.gemm(f"{B}/attn/Gemm")) _mlp(r, params, f"{B}/mlp1", f"{BN}.mlp1", "tanh") q_w = np.asarray(r.w.matmul_weight(f"{B}/cross_attn/MatMul"), np.float32) _put_linear(params, f"{BN}.cross_attn.q", (q_w, r.bias_after(f"{B}/cross_attn/MatMul"))) # K|V of this block: the bias is the constant input of cross_attn/Add_5, the weight a column block of the # unnamed [256, 1536] MatMul whose output is split three ways (one 512-wide K|V slice per block) add5 = r.w.node(f"{B}/cross_attn/Add_5") kv_b = r.const_input(add5.name) (kv_in,) = [t for t in add5.inputs if t and not r.w.has(t)] split = r.w.producer(kv_in) if split is None or split.op_type != "Split": raise ValueError(f"{B}/cross_attn/Add_5: K|V does not come from a Split") part = list(split.outputs).index(kv_in) sizes = [int(s) for s in np.asarray(r.w.array(split.inputs[1])).reshape(-1)] fused = r.w.producer(split.inputs[0]) fused_w = np.asarray(r.w.matmul_weight(fused.name), np.float32) start = int(sum(sizes[:part])) _put_linear(params, f"{BN}.cross_attn.kv", (fused_w[:, start:start + sizes[part]].copy(), kv_b)) _put_linear(params, f"{BN}.cross_attn.out", r.gemm(f"{B}/cross_attn/Gemm")) _mlp(r, params, f"{B}/mlp2", f"{BN}.mlp2", "tanh") scales.add(r.scalar(f"{B}/attn/Mul_2")) scales.add(r.scalar(f"{B}/cross_attn/Mul_2")) fills.add(float(np.asarray(r.w.param("/dit/blocks.0/attn/Where", 1)).reshape(-1)[0])) # shared by the blocks F = "/dit/final_layer" _put_linear(params, "decoder.dit.final_layer.adaLN_modulation", r.linear(f"{F}/adaLN_modulation/adaLN_modulation.1")) _put_ln(params, "decoder.dit.final_layer.norm_final", r.layer_norm(f"{F}/norm_final")) _put_ln(params, "decoder.dit.final_layer.proj.0", r.layer_norm(f"{F}/proj/proj.0")) _put_linear(params, "decoder.dit.final_layer.proj.1", r.linear(f"{F}/proj/proj.1")) gelu = r.gelu(f"{F}/proj/proj.2") if gelu != "tanh": raise ValueError(f"{F}/proj/proj.2: GELU approximate={gelu!r}, expected 'tanh'") _put_ln(params, "decoder.dit.final_layer.proj.3", r.layer_norm(f"{F}/proj/proj.3")) _put_linear(params, "decoder.dit.final_layer.proj.4", r.linear(f"{F}/proj/proj.4")) return {"gelu": r.facts["gelu"], "ln_eps": r.facts["ln_eps"], "attn_scale": scales, "mask_fill": fills, "sha256": r.w.sha256, "nodes": len(r.w.nodes())} def _read_turn(path: Path, params: Dict[str, np.ndarray]) -> Dict[str, Any]: r = _Reader(path) _put_linear(params, "decoder.turn_indicator_predictor", r.linear("/turn_indicator_predictor")) return {"sha256": r.w.sha256, "nodes": len(r.w.nodes())} EXPECTED_SHAPES = { # spot checks of the canonical layout (SPEC 4.2-4.4) "encoder.neighbor_encoder.channel_pre_project.fc1.w": (C.NEIGHBOR_FEATURE_DIM, C.MIXER_CHANNELS), "encoder.neighbor_encoder.token_pre_project.fc1.w": (C.INPUT_T + 1, C.MIXER_TOKENS), "encoder.lane_encoder.token_pre_project.fc1.w": (C.POINTS_PER_SEGMENT, C.MIXER_TOKENS), "encoder.polygon_encoder.channel_pre_project.fc1.w": (C.POLYGON_FEATURE_DIM, C.MIXER_CHANNELS), "encoder.polygon_encoder.token_pre_project.fc1.w": (C.POINTS_PER_POLYGON, C.MIXER_TOKENS), "encoder.line_string_encoder.channel_pre_project.fc1.w": (C.LINE_STRING_FEATURE_DIM, C.MIXER_CHANNELS), "encoder.lane_encoder.attribute_emb.w": (C.LANE_ATTRIBUTE_DIM, C.MIXER_CHANNELS), "encoder.turn_indicator_encoder.channel_pre_project.fc1.w": (C.TURN_INDICATOR_HISTORY, C.MIXER_CHANNELS), "encoder.pos_emb.w": (C.POS_FEATURE_DIM, C.HIDDEN_DIM), "encoder.fusion.blocks.0.attn.kv.w": (C.HIDDEN_DIM, 2 * C.HIDDEN_DIM), "decoder.dit.preproj.fc1.w": (C.DIT_INPUT_DIM, 512), "decoder.dit.t_embedder.fc1.w": (C.DIT_TIME_DIM, 512), "decoder.dit.blocks.0.adaLN_modulation.w": (C.HIDDEN_DIM, 6 * C.HIDDEN_DIM), "decoder.dit.blocks.0.attn.qkv.w": (C.HIDDEN_DIM, 3 * C.HIDDEN_DIM), "decoder.dit.blocks.2.cross_attn.kv.w": (C.HIDDEN_DIM, 2 * C.HIDDEN_DIM), "decoder.dit.final_layer.proj.4.w": (C.DIT_MLP_DIM, C.DIT_INPUT_DIM), "decoder.turn_indicator_predictor.w": (2 * len(C.TURN_HEAD_STEPS) + C.HIDDEN_DIM, C.TURN_INDICATOR_OUTPUT_DIM), } def _check_facts(facts: Dict[str, Any]) -> None: eps = facts["encoder"]["ln_eps"] | facts["decoder"]["ln_eps"] if eps != {C.LN_EPS} and not all(math.isclose(e, C.LN_EPS, rel_tol=1e-6) for e in eps): raise ValueError(f"LayerNorm epsilons {sorted(eps)}, expected {C.LN_EPS}") scales = facts["encoder"]["attn_scale"] | facts["decoder"]["attn_scale"] if len(scales) != 1 or not math.isclose(scales.pop(), 1.0 / math.sqrt(C.HEAD_DIM), rel_tol=1e-6): raise ValueError(f"attention scales {facts['encoder']['attn_scale'] | facts['decoder']['attn_scale']}") fills = facts["encoder"]["mask_fill"] | facts["decoder"]["mask_fill"] if fills != {float("-inf")}: raise ValueError(f"attention mask fill values {fills}, expected -inf") def find_weights_dir(explicit: Optional[str] = None) -> Optional[Path]: """A local directory holding the v5.0 files, or None: ``explicit`` > ``$DIFFUSION_PLANNER_WEIGHTS_DIR`` > the workspace download (``assets/diffusion-planner/hf_diffusion_planner``) > the HF cache snapshot of the pinned revision (``local_files_only``; never a network access).""" import os names = (C.ENCODER_ONNX, C.DECODER_ONNX, C.TURN_INDICATOR_ONNX, C.PARAM_JSON) cands: List[Path] = [] for c in (explicit, os.environ.get("DIFFUSION_PLANNER_WEIGHTS_DIR")): if c: cands.append(Path(c).expanduser()) here = Path(__file__).resolve() for parent in here.parents: cands.append(parent / "assets" / "diffusion-planner" / "hf_diffusion_planner") for c in cands: if all((c / n).is_file() for n in names): return c try: # the HF cache of `from_pretrained` (offline lookup only) from huggingface_hub import snapshot_download p = Path(snapshot_download("AutowareFoundation/diffusion_planner", revision="423efde67f5414734da43a7ad856c17ceb8b51aa", allow_patterns=list(names), local_files_only=True)) if all((p / n).is_file() for n in names): return p except Exception: # noqa: BLE001 -- not cached / no huggingface_hub: no weights pass return None def load_weights(weights_dir: Path, *, verify_sha256: bool = True) -> PlannerWeights: """Read the encoder / decoder / turn-indicator ONNX files and the param JSON of ``weights_dir``.""" weights_dir = Path(weights_dir) sha = {n: file_sha256(weights_dir / n) for n in C.FILE_SHA256} if verify_sha256: bad = {n: s for n, s in sha.items() if s != C.FILE_SHA256[n]} if bad: raise ValueError(f"{weights_dir}: files differ from AutowareFoundation/diffusion_planner@v5.0: " f"{sorted(bad)} (pass verify_sha256=False to load another export)") params: Dict[str, np.ndarray] = {} facts = {"encoder": _read_encoder(weights_dir / C.ENCODER_ONNX, params), "decoder": _read_decoder(weights_dir / C.DECODER_ONNX, params), "turn": _read_turn(weights_dir / C.TURN_INDICATOR_ONNX, params)} _check_facts(facts) for name, shape in EXPECTED_SHAPES.items(): if tuple(params[name].shape) != shape: raise ValueError(f"{name}: shape {params[name].shape}, expected {shape}") for name, arr in params.items(): if arr.dtype != np.float32 or not np.isfinite(arr).all(): raise ValueError(f"{name}: not finite float32") norm = load_param_json(weights_dir / C.PARAM_JSON) return PlannerWeights(params, norm, sha, facts, weights_dir) def param_count(params: Mapping[str, np.ndarray]) -> int: return int(sum(v.size for v in params.values())) def coverage(weights: PlannerWeights) -> Dict[str, Any]: """Proof that the canonical dict is a re-labelling of the export: every float initializer (size > 1) of the three files equals one canonical tensor (or the concatenation it was split from: the fused cross K/V MatMul, the two agent-embedding rows), and no canonical tensor holds the same data twice except the biases the export itself deduplicated. Returns ``{"unused": [...], "duplicates": [...], "initializers": n}`` (tests require both lists to be empty / the known pair).""" import hashlib import onnx from onnx import numpy_helper def h(a: np.ndarray) -> str: a = np.ascontiguousarray(np.asarray(a, np.float32)) return hashlib.sha1(a.tobytes() + str(a.shape).encode()).hexdigest() p = weights.params known = {h(v): k for k, v in p.items()} kv = np.concatenate([p[f"decoder.dit.blocks.{i}.cross_attn.kv.w"] for i in range(C.DIT_DEPTH)], axis=1) known[h(kv)] = "decoder.dit.blocks.*.cross_attn.kv.w (fused)" emb = p["decoder.dit.agent_embedding"] known[h(emb[0:1])] = "decoder.dit.agent_embedding[0]" known[h(emb[1:2])] = "decoder.dit.agent_embedding[1]" rpe = p["encoder.route_position_embedding"] known[h(rpe.reshape(1, *rpe.shape))] = "encoder.route_position_embedding" for mod in ("lane_encoder", "route_encoder"): unk = p[f"encoder.{mod}.unknown_speed_emb"] known[h(unk.reshape(1, -1))] = f"encoder.{mod}.unknown_speed_emb" for k, v in p.items(): # Gemm weights are stored [out, in] if k.endswith(".w"): known.setdefault(h(v.T), k) unused, total = [], 0 for f in (C.ENCODER_ONNX, C.DECODER_ONNX, C.TURN_INDICATOR_ONNX): for init in onnx.load(str(weights.path / f)).graph.initializer: a = numpy_helper.to_array(init) if a.dtype != np.float32 or a.size <= 1: continue total += 1 if h(a) not in known: unused.append(f"{f}:{init.name}{list(a.shape)}") seen: Dict[str, List[str]] = {} for k, v in p.items(): seen.setdefault(h(v), []).append(k) dups = sorted(tuple(sorted(ks)) for ks in seen.values() if len(ks) > 1) return {"unused": unused, "duplicates": dups, "initializers": total}