File size: 10,858 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 | # 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
|