# 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 --------------------------------------------------------------------------------------------- @dataclass 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)) @torch.no_grad() 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 --------------------------------------------------------------------------------------- @dataclass 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