Download code/tt_diffusion_planner/reference/rewrites.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 10.9 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/reference/rewrites.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/reference/rewrites.py
-
curl -L -o rewrites.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/reference/rewrites.py
10.9 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """Exact rewrites of the deployed graph that the TT port applies, as constants and CPU forwards (SPEC 4.6, PLAN 2.12). | |
| Each rewrite is exact in real arithmetic; its constants are computed in float64 from the float32 weights and rounded | |
| once to float32 (the "fold in fp64, round once" rule of ``ttaw.weights``). The CPU forwards here use them and are | |
| tested against the as-exported :mod:`.model` (``tests/test_reference_host.py``), so the device port has a CPU oracle in | |
| its own parameterisation. | |
| 1. **Per-step adaLN tables folded into the LayerNorm affine** (SPEC 4.6.2, 8.1). With a uniform diffusion time (the | |
| multi-step mode) the t-embedding and every adaLN output are the same for all 321 agents and depend only on t, so | |
| for the 11 evaluation times ``modulate(LN(x; g, b), shift, scale) = LN(x; g (1 + scale), b (1 + scale) + shift)`` | |
| and the gates are per-step ``[256]`` rows. :func:`adaln_tables` builds them for ``SolverPlan.eval_times``. | |
| 2. **Cross-attention K/V hoisted** (SPEC 4.6.1): ``encoding @ W_kv + b_kv`` of the three blocks once per plan | |
| (the decoder graph recomputes it at each of the 11 calls). | |
| 3. **Pad-relative fp32 pre-projection island** (SPEC 4.6.7, probe P12): after the in-graph history truncation the | |
| ego (rows 6..30) and neighbour (rows 0..24) inputs of ``channel_pre_project`` are zero rows, which all map to | |
| ``c0 = fc2(gelu(b1)) + b2``. With the all-zero ("pad") agent's outputs ``t1_pad = b_t1 + c0 (x) sum_t W_t1[t]``, | |
| ``g_pad = gelu(t1_pad)``, ``t2_pad = g_pad @ W_t2 + b_t2`` the island becomes | |
| ``t1 = t1_pad + (z_valid - c0)^T @ W_t1[valid rows]`` and ``t2 = t2_pad + (gelu(t1) - g_pad) @ W_t2``, so the | |
| device's TF32-like fp32 matmuls only see deviations from the pad agent (P12: deviation PCC 0.995706 -> 0.999991 | |
| on ``straight``). ``gelu_b1 = gelu(b_c1)`` also lets ``channel_pre`` itself run pad-relative: | |
| ``z - c0 = (gelu(x @ W_c1 + b_c1) - gelu_b1) @ W_c2``. | |
| """ | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from typing import Dict, Mapping, Sequence, Tuple | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| from . import config as C | |
| from .model import Decoder, _t | |
| __all__ = ["AdaLNTables", "adaln_tables", "decoder_forward_folded", "IslandConstants", "island_constants", | |
| "island_forward_pad_relative", "island_forward_exported", "cross_kv"] | |
| # ---- 1. adaLN tables --------------------------------------------------------------------------------------------- | |
| class AdaLNTables: | |
| """float32 tables for K evaluation times: ``temb [K, 256]``; per DiT block ``i``: ``norm1_gamma``, | |
| ``norm1_beta``, ``gate_msa``, ``norm2_gamma``, ``norm2_beta``, ``gate_mlp`` ``[K, 256]``; the final layer's | |
| ``final_gamma`` / ``final_beta`` ``[K, 256]``.""" | |
| eval_times: Tuple[float, ...] | |
| temb: np.ndarray | |
| blocks: Tuple[Dict[str, np.ndarray], ...] | |
| final_gamma: np.ndarray | |
| final_beta: np.ndarray | |
| def adaln_tables(params: Mapping[str, np.ndarray], eval_times: Sequence[float]) -> AdaLNTables: | |
| p = {k: torch.from_numpy(np.array(v, np.float64)) for k, v in params.items() if k.startswith("decoder.dit")} | |
| def lin(x, n): | |
| return x @ p[f"{n}.w"] + p[f"{n}.b"] | |
| t = torch.tensor([[float(np.float32(v))] * C.DIT_TIME_DIM for v in eval_times], dtype=torch.float64) | |
| c = lin(F.gelu(lin(t, "decoder.dit.t_embedder.fc1")), "decoder.dit.t_embedder.fc2") # [K, 256] | |
| silu = c * torch.sigmoid(c) | |
| blocks = [] | |
| for i in range(C.DIT_DEPTH): | |
| B = f"decoder.dit.blocks.{i}" | |
| sh_msa, sc_msa, g_msa, sh_mlp, sc_mlp, g_mlp = torch.split(lin(silu, f"{B}.adaLN_modulation"), C.HIDDEN_DIM, -1) | |
| g1, b1 = p[f"{B}.norm1.gamma"], p[f"{B}.norm1.beta"] | |
| g2, b2 = p[f"{B}.norm2.gamma"], p[f"{B}.norm2.beta"] | |
| blocks.append({k: v.to(torch.float32).numpy() for k, v in { | |
| "norm1_gamma": g1 * (1 + sc_msa), "norm1_beta": b1 * (1 + sc_msa) + sh_msa, "gate_msa": g_msa, | |
| "norm2_gamma": g2 * (1 + sc_mlp), "norm2_beta": b2 * (1 + sc_mlp) + sh_mlp, "gate_mlp": g_mlp}.items()}) | |
| Fn = "decoder.dit.final_layer" | |
| sh, sc = torch.split(lin(silu, f"{Fn}.adaLN_modulation"), C.HIDDEN_DIM, -1) | |
| gf, bf = p[f"{Fn}.norm_final.gamma"], p[f"{Fn}.norm_final.beta"] | |
| return AdaLNTables(tuple(float(v) for v in eval_times), c.to(torch.float32).numpy(), tuple(blocks), | |
| (gf * (1 + sc)).to(torch.float32).numpy(), (bf * (1 + sc) + sh).to(torch.float32).numpy()) | |
| def cross_kv(params: Mapping[str, np.ndarray], encoding: np.ndarray) -> Tuple[np.ndarray, ...]: | |
| """Hoisted cross-attention K|V ``[564, 512]`` of each DiT block (float32 matmul, as the exported graph).""" | |
| enc = np.asarray(encoding, np.float32) | |
| return tuple((enc @ params[f"decoder.dit.blocks.{i}.cross_attn.kv.w"] | |
| + params[f"decoder.dit.blocks.{i}.cross_attn.kv.b"]).astype(np.float32) for i in range(C.DIT_DEPTH)) | |
| def decoder_forward_folded(dec: Decoder, tables: AdaLNTables, k: int, x_t: np.ndarray, kv: Sequence[torch.Tensor], | |
| agent_valid: np.ndarray) -> torch.Tensor: | |
| """One decoder evaluation with the step-``k`` tables: LayerNorms with folded affine (no SiLU / adaLN matmuls), | |
| gates as constant rows. Same result as ``Decoder.forward`` up to float32 rounding.""" | |
| from .model import _key_bias | |
| P = int(np.shape(x_t)[0]) | |
| x = dec.mlp(_t(x_t).reshape(P, C.DIT_INPUT_DIM), "decoder.dit.preproj") | |
| emb = dec.p["decoder.dit.agent_embedding"] | |
| x = x + torch.cat([emb[0:1], emb[1:2].expand(P - 1, -1)], dim=0) | |
| key_bias = _key_bias(agent_valid) | |
| def ln_folded(h, gamma, beta): | |
| return F.layer_norm(h, h.shape[-1:], _t(gamma), _t(beta), eps=C.LN_EPS) | |
| for i in range(C.DIT_DEPTH): | |
| B, T = f"decoder.dit.blocks.{i}", tables.blocks[i] | |
| h = ln_folded(x, T["norm1_gamma"][k], T["norm1_beta"][k]) | |
| q, kk, v = torch.split(dec.linear(h, f"{B}.attn.qkv"), C.HIDDEN_DIM, dim=-1) | |
| x = x + _t(T["gate_msa"][k]) * dec.linear(dec.attend(q, kk, v, key_bias), f"{B}.attn.out") | |
| h = ln_folded(x, T["norm2_gamma"][k], T["norm2_beta"][k]) | |
| x = x + _t(T["gate_mlp"][k]) * dec.mlp(h, f"{B}.mlp1", "tanh") | |
| heads = dec.attention(dec.ln(x, f"{B}.norm3"), kv[i], f"{B}.cross_attn.q", None) | |
| x = x + dec.linear(heads, f"{B}.cross_attn.out") | |
| x = x + dec.mlp(dec.ln(x, f"{B}.norm4"), f"{B}.mlp2", "tanh") | |
| Fn = "decoder.dit.final_layer" | |
| h = ln_folded(x, tables.final_gamma[k], tables.final_beta[k]) | |
| h = F.gelu(dec.linear(dec.ln(h, f"{Fn}.proj.0"), f"{Fn}.proj.1"), approximate="tanh") | |
| return dec.linear(dec.ln(h, f"{Fn}.proj.3"), f"{Fn}.proj.4").reshape(P, C.OUTPUT_T + 1, C.POSE_DIM) | |
| # ---- 3. pad-relative island --------------------------------------------------------------------------------------- | |
| class IslandConstants: | |
| """Pad-agent constants of one pre-projection island (float32, built in float64). ``valid_rows`` are the time | |
| rows that can be non-zero after the truncation (ego 0..5, neighbours 25..30); ``w_t1_valid`` is | |
| ``W_t1[valid_rows]`` ``[6, 64]``.""" | |
| category: str | |
| valid_rows: Tuple[int, ...] | |
| gelu_b1: np.ndarray # [128] gelu(b_c1): channel_pre's first activation of a zero row | |
| c0: np.ndarray # [128] channel_pre of a zero row | |
| t1_pad: np.ndarray # [128, 64] | |
| g_pad: np.ndarray # [128, 64] | |
| t2_pad: np.ndarray # [128, 64] | |
| w_t1_valid: np.ndarray # [6, 64] | |
| ISLAND_ROWS = {"neighbor": tuple(range(C.NEIGHBOR_HISTORY_KEEP.start, C.NEIGHBOR_HISTORY_KEEP.stop)), | |
| "ego": tuple(range(C.EGO_HISTORY_KEEP.start, C.EGO_HISTORY_KEEP.stop))} | |
| ISLAND_MODULE = {"neighbor": "encoder.neighbor_encoder", "ego": "encoder.ego_encoder"} | |
| def island_constants(params: Mapping[str, np.ndarray], category: str) -> IslandConstants: | |
| N = ISLAND_MODULE[category] | |
| w = {k: torch.from_numpy(np.array(v, np.float64)) for k, v in params.items() if k.startswith(N + ".")} | |
| gelu_b1 = F.gelu(w[f"{N}.channel_pre_project.fc1.b"]) | |
| c0 = gelu_b1 @ w[f"{N}.channel_pre_project.fc2.w"] + w[f"{N}.channel_pre_project.fc2.b"] | |
| w_t1 = w[f"{N}.token_pre_project.fc1.w"] # [31, 64] | |
| t1_pad = w[f"{N}.token_pre_project.fc1.b"][None, :] + c0[:, None] * w_t1.sum(0)[None, :] | |
| g_pad = F.gelu(t1_pad) | |
| t2_pad = g_pad @ w[f"{N}.token_pre_project.fc2.w"] + w[f"{N}.token_pre_project.fc2.b"] | |
| rows = ISLAND_ROWS[category] | |
| f32 = lambda t: t.to(torch.float32).numpy() # noqa: E731 | |
| return IslandConstants(category, rows, f32(gelu_b1), f32(c0), f32(t1_pad), f32(g_pad), f32(t2_pad), | |
| f32(w_t1[list(rows)])) | |
| def island_forward_exported(params: Mapping[str, np.ndarray], category: str, x: np.ndarray, | |
| dtype: torch.dtype = torch.float64) -> Tuple[torch.Tensor, torch.Tensor]: | |
| """The island as exported: ``x`` ``[E, 31, C_in]`` -> ``t1``, ``t2`` ``[E, 128, 64]`` (``dtype`` math).""" | |
| N = ISLAND_MODULE[category] | |
| w = {k: torch.from_numpy(np.array(v, np.float64)).to(dtype) for k, v in params.items() if k.startswith(N + ".")} | |
| def lin(h, n): | |
| return h @ w[f"{N}.{n}.w"] + w[f"{N}.{n}.b"] | |
| z = lin(F.gelu(lin(torch.from_numpy(np.asarray(x, np.float64)).to(dtype), "channel_pre_project.fc1")), | |
| "channel_pre_project.fc2") | |
| t1 = lin(z.transpose(1, 2), "token_pre_project.fc1") | |
| return t1, lin(F.gelu(t1), "token_pre_project.fc2") | |
| def island_forward_pad_relative(params: Mapping[str, np.ndarray], consts: IslandConstants, x: np.ndarray, | |
| dtype: torch.dtype = torch.float64) -> Tuple[torch.Tensor, torch.Tensor]: | |
| """The pad-relative island on the valid rows only (``x`` ``[E, 31, C_in]``; rows outside ``valid_rows`` must be | |
| zero, which the truncation guarantees). Matmuls see deviations from the pad agent only.""" | |
| N = ISLAND_MODULE[consts.category] | |
| w = {k: torch.from_numpy(np.array(v, np.float64)).to(dtype) for k, v in params.items() if k.startswith(N + ".")} | |
| cst = {k: torch.from_numpy(np.array(getattr(consts, k), np.float64)).to(dtype) | |
| for k in ("gelu_b1", "c0", "t1_pad", "g_pad", "t2_pad", "w_t1_valid")} | |
| xv = torch.from_numpy(np.asarray(x, np.float64)[:, list(consts.valid_rows)]).to(dtype) # [E, 6, C_in] | |
| h = F.gelu(xv @ w[f"{N}.channel_pre_project.fc1.w"] + w[f"{N}.channel_pre_project.fc1.b"]) - cst["gelu_b1"] | |
| dz = h @ w[f"{N}.channel_pre_project.fc2.w"] # z - c0, [E, 6, 128] | |
| t1 = cst["t1_pad"] + dz.transpose(1, 2) @ cst["w_t1_valid"] # [E, 128, 64] | |
| t2 = cst["t2_pad"] + (F.gelu(t1) - cst["g_pad"]) @ w[f"{N}.token_pre_project.fc2.w"] | |
| return t1, t2 | |