Files
dela 49aede9cb2 Fit 0.5b training on 32GB: SDPA MLA, block checkpoint, chunked CE
Whole-mixer checkpoint plus T×T MLA scores OOM'd a 31GB GPU on backward.
Checkpoint each AttnRes block, run absorbed MLA through SDPA, and compute
CE in vocab chunks so [B,T,V] logits are never materialized.

--max-tokens is now the training budget; default --steps 2000 no longer
caps a 1B-token run at 250 optimizer steps.
2026-08-25 20:09:27 +08:00

533 lines
17 KiB
Python

"""
Attention Residual in one file
Reference:
Kimi Team, Guangyu Chen, Yu Zhang, Jianlin Su, Weixin Xu, Siyuan Pan,
Yaoyu Wang, Yucheng Wang, Guanduo Chen, et al.
"Attention Residuals." arXiv:2603.15031, 2026.
https://arxiv.org/abs/2603.15031
This module is a compact PyTorch reference implementation of:
- Full AttnRes
- Block AttnRes
- two-phase inter/intra-block computation from the paper
CausalLM wires Full/Block stacks when ``config.attnres`` is ``full`` or
``block``. Standard residual (``x += attn; x += ffn``) is ``attnres="off"``.
"""
import torch
import torch.nn.functional as F
from einops import rearrange
from torch import Tensor, nn
from torch.utils.checkpoint import checkpoint as activation_checkpoint
ATTNRES_MODES = ("off", "full", "block")
def exists(x):
return x is not None
def validate_attnres(mode: str, block_size: int | None) -> None:
if mode not in ATTNRES_MODES:
raise ValueError(f"attnres must be one of {ATTNRES_MODES}, got {mode!r}")
if block_size is not None and block_size < 1:
raise ValueError(f"attnres_block_size must be >= 1, got {block_size}")
def atomic_block_size(num_hidden_layers: int, attnres_block_size: int | None) -> int:
"""DecoderBlocks per AttnRes block, converted to attn|ffn atomic layers.
``None`` targets about 8 blocks: ``max(1, ceil(L / 8))`` DecoderBlocks.
"""
layers_per_block = (
attnres_block_size
if attnres_block_size is not None
else max(1, (num_hidden_layers + 7) // 8)
)
if layers_per_block < 1:
raise ValueError(f"attnres_block_size must be >= 1, got {layers_per_block}")
return layers_per_block * 2
class BorrowedSubLayer(nn.Module):
"""``fn(norm(x))`` without registering ``norm``/``fn`` (owned by DecoderBlock)."""
def __init__(self, norm: nn.Module, fn: nn.Module):
super().__init__()
self._borrowed = (norm, fn)
def forward(self, x: Tensor) -> Tensor:
norm, fn = self._borrowed
return fn(norm(x))
def rms(x: Tensor, eps: float):
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + eps)
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x: Tensor) -> Tensor:
return rms(x, self.eps) * self.weight
class DepthResidual(nn.Module):
"""
h_l = sum_i softmax_i(w_l^T RMSNorm(v_i))*v_i
Keep query and RMSNorm gain separate
Since q^T (gamma * RMS(v)) == (q * gamma)^T RMS(v),
we can fold gamma into q for scoring.
"""
def __init__(self, dim: int, eps: float = 1e-8, zero_init: bool = True):
super().__init__()
self.query = nn.Parameter(torch.zeros(dim))
self.norm = RMSNorm(dim, eps=eps)
if not zero_init:
nn.init.normal_(self.query, std=0.02)
def effective_query(self) -> Tensor:
return (self.query * self.norm.weight).float()
def logits(self, sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
sources = stack_layers(sources) # [n, b, t, d]
q = self.effective_query() # [d]
k = rms(sources.float(), self.norm.eps) # [n, b, t, d]
return torch.einsum("d, n b t d -> n b t", q, k)
def forward(self, sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
sources = stack_layers(sources)
weights = self.logits(sources).softmax(dim=0)
out = torch.einsum("n b t, n b t d -> b t d", weights, sources.float())
return out.to(sources.dtype)
class DepthResidualList(nn.Module):
def __init__(self, dim: int, depth: int, eps: float, zero_init: bool = True):
super().__init__()
# for L layers (depth), create depth residual modules
self.layers = nn.ModuleList(
[DepthResidual(dim, eps=eps, zero_init=zero_init) for _ in range(depth)]
)
def __getitem__(self, idx: int) -> DepthResidual:
return self.layers[idx]
def __iter__(self):
return iter(self.layers)
def __len__(self):
return len(self.layers)
# attnres stacks
class FullAttnResStack(nn.Module):
"""
Full AttnRes over atomic layers
eg: f_1,...,f_L
Each entry in `layers` should already be a full atomic layer fxn
x -> f_l(x)
"""
def __init__(
self,
dim: int,
layers,
*,
eps: float = 1e-8,
zero_init_queries: bool = True,
is_final_aggregate: bool = True,
):
super().__init__()
self.layers = nn.ModuleList(list(layers))
self.eps = eps
depth = len(self.layers)
self.residuals = DepthResidualList(dim, depth, eps, zero_init_queries)
self.final_residual = (
DepthResidual(dim, eps, zero_init_queries) if is_final_aggregate else None
)
def forward_naive(self, x: Tensor) -> Tensor:
sources = [x]
for layer, residual in zip(self.layers, self.residuals):
h = residual(sources)
out = layer(h)
sources.append(out)
return (
self.final_residual(sources) if exists(self.final_residual) else sources[-1]
)
def forward_two_phase(self, x: Tensor, schedule_block_size: int) -> Tensor:
assert schedule_block_size > 0
sources = [x]
depth = len(self.layers)
start = 0
while start < depth:
end = min(start + schedule_block_size, depth)
queries = torch.stack(
[self.residuals[i].effective_query() for i in range(start, end)], dim=0
)
inter_sources = stack_layers(sources)
inter_stats = attn_with_stats(queries, inter_sources, self.eps)
local_outputs = [] # outputs of intra-block
for local_idx, layer_idx in enumerate(range(start, end)):
stats = inter_stats.select(local_idx)
if len(local_outputs) > 0:
intra_sources = stack_layers(local_outputs)
intra = attn_with_stats(
queries[local_idx : local_idx + 1], intra_sources, self.eps
).select(0)
stats = merge_attn_stats(stats, intra)
h = stats.normalized()
out = self.layers[layer_idx](h)
local_outputs.append(out)
sources.append(out)
start = end
return (
self.final_residual(sources) if exists(self.final_residual) else sources[-1]
)
def forward(self, x: Tensor, schedule_block_size: int | None = None) -> Tensor:
if schedule_block_size is None:
return self.forward_naive(x)
return self.forward_two_phase(x, schedule_block_size)
class BlockAttnResStack(nn.Module):
"""
Block AttnRes over atomic layers
`block_size` is in atomic layers, not Transformer blocks.
Eg: block_size=8 -> 4 transformer blocks when layers alternate attn/MLP
The default forward path is the two-phase algorithm from the paper:
phase 1: batch inter-block attn from all queries in the block
phase 2: merge the evolving intra-block partial sum with online softmax
"""
def __init__(
self,
dim: int,
layers,
*,
block_size: int,
eps: float = 1e-8,
zero_init_queries: bool = True,
is_final_aggregate: bool = True,
):
super().__init__()
self.layers = nn.ModuleList(list(layers))
assert len(self.layers) > 0
assert block_size > 0
self.block_size = block_size
self.eps = eps
depth = len(self.layers)
self.residuals = DepthResidualList(
dim, depth, eps=eps, zero_init=zero_init_queries
)
self.final_residual = (
DepthResidual(dim, eps=eps, zero_init=zero_init_queries)
if is_final_aggregate
else None
)
self.gradient_checkpointing = False
def forward_naive(self, x: Tensor) -> Tensor:
blocks = [x] # b_0=embedding/input representation
partial = None
for layer_idx, (layer, residual) in enumerate(
zip(self.layers, self.residuals), start=1
):
sources = blocks if partial is None else blocks + [partial]
h = residual(sources)
out = layer(h)
partial = out if partial is None else (partial + out)
if (layer_idx % self.block_size == 0) or (layer_idx == len(self.layers)):
blocks.append(partial)
partial = None
return (
self.final_residual(blocks) if exists(self.final_residual) else blocks[-1]
)
def _run_block_two_phase(
self, blocks: list[Tensor], start: int, end: int
) -> Tensor:
queries = torch.stack(
[self.residuals[i].effective_query() for i in range(start, end)], dim=0
)
inter_sources = stack_layers(blocks)
inter = attn_with_stats(queries, inter_sources, self.eps)
partial = None
for local_idx, layer_idx in enumerate(range(start, end)):
stats = inter.select(local_idx)
if partial is not None:
intra = single_source_stats(queries[local_idx], partial, self.eps)
stats = merge_attn_stats(stats, intra)
h = stats.normalized()
out = self.layers[layer_idx](h)
partial = out if partial is None else (partial + out)
return partial
def forward(self, x: Tensor) -> Tensor:
blocks = [x]
depth = len(self.layers)
start = 0
use_ckpt = (
self.gradient_checkpointing and self.training and torch.is_grad_enabled()
)
while start < depth:
end = min(start + self.block_size, depth)
if use_ckpt:
def _run(*srcs, _start=start, _end=end):
return self._run_block_two_phase(list(srcs), _start, _end)
blocks.append(
activation_checkpoint(_run, *blocks, use_reentrant=False)
)
else:
blocks.append(self._run_block_two_phase(blocks, start, end))
start = end
return (
self.final_residual(blocks) if exists(self.final_residual) else blocks[-1]
)
# helpers
def stack_layers(sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
if isinstance(sources, Tensor):
assert sources.ndim == 4, f"expected [n, b, t, d] got {tuple(sources.shape)}"
return sources
assert len(sources) > 0, "needs at least one source"
return torch.stack(tuple(sources), dim=0)
class SingleAttnStats:
def __init__(self, numer: Tensor, denom: Tensor, max: Tensor):
self.numer = numer # [b,t,d]
self.max = max # [b,t]
self.denom = denom # [b,t]
def normalized(self) -> Tensor:
return self.numer / self.denom[..., None]
class AttnStats:
# store the numerator => e^{s_{j}-m} * v_j where m is the max score so far
# store the max m = max(s_j)
# store the denominator sum_j e^{s_{j}-m}
def __init__(self, numer: Tensor, denom: Tensor, max: Tensor):
self.numer = numer # [q,b,t,d]
self.max = max # [q,b,t]
self.denom = denom # [q,b,t]
def select(self, idx: int) -> "SingleAttnStats":
return SingleAttnStats(self.numer[idx], self.denom[idx], self.max[idx])
def attn_with_stats(queries: Tensor, sources: Tensor, eps: float = 1e-8) -> AttnStats:
"""
queries: [q, d]
sources: [n, b, t, d]
Returns the following for online softmax:
numer = sum_i exp(logit_i - m)*v_i
m = max_i logit_i
denom = sum_i exp(logit_i - m)
"""
normed = rms(sources, eps)
logits = torch.einsum("q d, n b t d -> q n b t", queries, normed)
m = logits.amax(dim=1)
weights = torch.exp(logits - m[:, None])
numer = torch.einsum("q n b t, n b t d -> q b t d", weights, sources)
denom = weights.sum(dim=1)
return AttnStats(numer, denom, m)
def single_source_stats(
query: Tensor, source: Tensor, eps: float = 1e-8
) -> SingleAttnStats:
score = torch.einsum("d, b t d -> b t", query, rms(source, eps))
denom = torch.ones_like(score)
return SingleAttnStats(source, denom, score)
def merge_attn_stats(a: SingleAttnStats, b: SingleAttnStats) -> SingleAttnStats:
m = torch.maximum(a.max, b.max)
wa = torch.exp(a.max - m)
wb = torch.exp(b.max - m)
numer = wa[..., None] * a.numer + wb[..., None] * b.numer
denom = wa * a.denom + wb * b.denom
return SingleAttnStats(numer, denom, m)
# transformer
class PreNorm(nn.Module):
def __init__(self, dim: int, fn: nn.Module, eps: float = 1e-8):
super().__init__()
self.norm = RMSNorm(dim, eps=eps)
self.fn = fn
def forward(self, x: Tensor) -> Tensor:
return self.fn(self.norm(x))
class CausalAttention(nn.Module):
def __init__(
self, dim: int, heads: int = 8, dim_head: int = 64, dropout: float = 0.0
):
super().__init__()
inner_dim = heads * dim_head
self.heads = heads
self.dim_head = dim_head
self.dropout = dropout
self.to_qkv = nn.Linear(dim, inner_dim * 3, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
def forward(self, x: Tensor) -> Tensor:
q, k, v = self.to_qkv(x).chunk(3, dim=-1)
def split_heads(y: Tensor) -> Tensor:
return rearrange(y, "b t (h d) -> b h t d", h=self.heads)
q, k, v = map(split_heads, (q, k, v))
out = F.scaled_dot_product_attention(
q, k, v, is_causal=True, dropout_p=self.dropout if self.training else 0.0
)
out = rearrange(out, "b h t d -> b t (h d)")
return self.to_out(out)
class SwiGLU(nn.Module):
def __init__(self, dim: int, mult: int = 4, dropout: float = 0.0):
# dropout not needed unless training on a smaller training data
super().__init__()
inner_dim = dim * mult
self.to_hidden = nn.Linear(dim, inner_dim * 2, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x: Tensor) -> Tensor:
gate, value = self.to_hidden(x).chunk(2, dim=-1)
x = F.silu(gate) * value
x = self.dropout(x)
return self.to_out(x)
class AttnResTransformer(nn.Module):
"""
Small GPT-style reference model using AttnRes
Using plain PyTorch: tok/pos embedding, alternating
causal attn, SwiGLU MLP layers, final norm, output head.
"""
def __init__(
self,
*,
num_tokens: int,
dim: int,
depth: int,
max_seq_len: int,
heads: int = 8,
dim_head: int = 64,
ff_mult: int = 4,
attn_dropout: float = 0.0,
ff_dropout: float = 0.0,
attnres: str = "block", # full or block
block_size: int = 8,
zero_init_queries: bool = True,
is_final_aggregate: bool = True,
eps: float = 1e-8,
):
super().__init__()
assert attnres in {"full", "block"}
self.max_seq_len = max_seq_len
self.attnres = attnres
self.token_emb = nn.Embedding(num_tokens, dim)
self.pos_emb = nn.Embedding(max_seq_len, dim)
atomic_layers = []
for _ in range(depth):
atomic_layers.append(
PreNorm(dim, CausalAttention(dim, heads, dim_head, attn_dropout), eps)
)
atomic_layers.append(PreNorm(dim, SwiGLU(dim, ff_mult, ff_dropout), eps))
if attnres == "full":
self.backbone = FullAttnResStack(
dim,
atomic_layers,
eps=eps,
zero_init_queries=zero_init_queries,
is_final_aggregate=is_final_aggregate,
)
else:
self.backbone = BlockAttnResStack(
dim,
atomic_layers,
block_size=block_size,
eps=eps,
zero_init_queries=zero_init_queries,
is_final_aggregate=is_final_aggregate,
)
self.final_norm = RMSNorm(dim, eps)
self.to_logits = nn.Linear(dim, num_tokens, bias=False)
def forward(self, ids: Tensor, schedule_block_size: int | None = None) -> Tensor:
b, t = ids.shape
assert t <= self.max_seq_len
pos = torch.arange(t, device=ids.device)
x = self.token_emb(ids) + self.pos_emb(pos)[None, :, :]
if self.attnres == "full":
x = self.backbone(x, schedule_block_size=schedule_block_size)
else:
x = self.backbone(x)
x = self.final_norm(x)
return self.to_logits(x)
__all__ = [
"ATTNRES_MODES",
"RMSNorm",
"DepthResidual",
"DepthResidualList",
"FullAttnResStack",
"BlockAttnResStack",
"BorrowedSubLayer",
"PreNorm",
"CausalAttention",
"SwiGLU",
"AttnResTransformer",
"atomic_block_size",
"validate_attnres",
]