changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
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 ---------------------------------------------------------------------------------------------
@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