File size: 24,407 Bytes
4d9b003
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
# 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}