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