Download code/tt_diffusion_planner/reference/weights.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 24.4 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/reference/weights.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/reference/weights.py
-
curl -L -o weights.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/reference/weights.py
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 | |
| 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 | |
| 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} | |