changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
24.4 kB
# 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: ``<module>.w`` ``[in, out]`` (``y = x @ w + b``) and ``<module>.b`` ``[out]``;
- LayerNorms: ``<module>.gamma`` / ``<module>.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.<lane|route>_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 ``<path>``: either ``<path>/MatMul``
followed by an ``Add`` (3-D inputs) or ``<path>/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}