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.
143 lines
5.0 KiB
Python
143 lines
5.0 KiB
Python
"""Causal LM stem: embed -> DecoderBlock* -> norm -> lm_head.
|
|
|
|
KDA-only and K3-like both use this class. Config.layer_specs() chooses
|
|
attn/ffn per layer: ("kda"|"mla", "swiglu"|"moe").
|
|
|
|
``config.attnres`` selects the depth mixer:
|
|
off — standard residual inside each DecoderBlock (default)
|
|
full — Full AttnRes over attn|ffn sublayers
|
|
block — Block AttnRes (K3); block size from ``attnres_block_size``
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import torch
|
|
import torch.nn.functional as F
|
|
from torch import nn
|
|
from torch.utils.checkpoint import checkpoint as activation_checkpoint
|
|
|
|
from ..layers.attn_res import (
|
|
BlockAttnResStack,
|
|
BorrowedSubLayer,
|
|
FullAttnResStack,
|
|
atomic_block_size,
|
|
)
|
|
from ..layers.block import DecoderBlock
|
|
from ..layers.rmsnorm import RMSNorm
|
|
|
|
|
|
def _chunked_linear_cross_entropy(
|
|
hidden: torch.Tensor,
|
|
weight: torch.Tensor,
|
|
labels: torch.Tensor,
|
|
ignore_index: int = -100,
|
|
chunk_size: int = 256,
|
|
) -> torch.Tensor:
|
|
"""CE without materializing [B, T, vocab]. Match mean reduction over valid labels."""
|
|
features = hidden[:, :-1].reshape(-1, hidden.size(-1))
|
|
targets = labels[:, 1:].reshape(-1)
|
|
total = hidden.new_zeros(())
|
|
n_valid = hidden.new_zeros((), dtype=torch.long)
|
|
for start in range(0, features.size(0), chunk_size):
|
|
sl = slice(start, start + chunk_size)
|
|
logits = F.linear(features[sl], weight)
|
|
total = total + F.cross_entropy(
|
|
logits, targets[sl], ignore_index=ignore_index, reduction="sum"
|
|
)
|
|
n_valid = n_valid + (targets[sl] != ignore_index).sum()
|
|
return total / n_valid.clamp_min(1).to(dtype=total.dtype)
|
|
|
|
|
|
def _build_mixer(config, blocks: nn.ModuleList):
|
|
mode = getattr(config, "attnres", "off")
|
|
if mode == "off":
|
|
return None
|
|
atomics = []
|
|
for block in blocks:
|
|
atomics.append(BorrowedSubLayer(block.attn_norm, block.attn))
|
|
atomics.append(BorrowedSubLayer(block.ffn_norm, block.ffn))
|
|
kwargs = dict(
|
|
eps=config.norm_eps,
|
|
zero_init_queries=getattr(config, "attnres_zero_init_queries", True),
|
|
is_final_aggregate=getattr(config, "attnres_final_aggregate", True),
|
|
)
|
|
if mode == "full":
|
|
return FullAttnResStack(config.hidden_size, atomics, **kwargs)
|
|
if mode == "block":
|
|
return BlockAttnResStack(
|
|
config.hidden_size,
|
|
atomics,
|
|
block_size=atomic_block_size(
|
|
config.num_hidden_layers, getattr(config, "attnres_block_size", None)
|
|
),
|
|
**kwargs,
|
|
)
|
|
raise ValueError(f"unknown attnres mode: {mode!r}")
|
|
|
|
|
|
class CausalLM(nn.Module):
|
|
def __init__(self, config):
|
|
super().__init__()
|
|
self.config = config
|
|
self.attnres = getattr(config, "attnres", "off")
|
|
self.embedding = nn.Embedding(config.vocab_size, config.hidden_size)
|
|
self.blocks = nn.ModuleList(
|
|
[
|
|
DecoderBlock.from_spec(config, attn, ffn)
|
|
for attn, ffn in config.layer_specs()
|
|
]
|
|
)
|
|
self.mixer = _build_mixer(config, self.blocks)
|
|
self.gradient_checkpointing = bool(
|
|
getattr(config, "gradient_checkpointing", False)
|
|
)
|
|
self.norm = RMSNorm(config.hidden_size, config.norm_eps)
|
|
self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
|
|
nn.init.normal_(self.embedding.weight, std=config.initializer_range)
|
|
nn.init.normal_(self.lm_head.weight, std=config.initializer_range)
|
|
if config.tie_word_embeddings:
|
|
self.lm_head.weight = self.embedding.weight
|
|
|
|
def forward(
|
|
self,
|
|
input_ids: torch.Tensor,
|
|
labels: torch.Tensor | None = None,
|
|
ignore_index: int = -100,
|
|
):
|
|
x = self.embedding(input_ids)
|
|
if self.mixer is None:
|
|
for block in self.blocks:
|
|
if self.gradient_checkpointing and self.training:
|
|
x = activation_checkpoint(block, x, use_reentrant=False)
|
|
else:
|
|
x = block(x)
|
|
else:
|
|
self.mixer.gradient_checkpointing = self.gradient_checkpointing
|
|
x = self.mixer(x)
|
|
hidden = self.norm(x)
|
|
if labels is None:
|
|
return self.lm_head(hidden)
|
|
return _chunked_linear_cross_entropy(
|
|
hidden, self.lm_head.weight, labels, ignore_index=ignore_index
|
|
)
|
|
|
|
@torch.inference_mode()
|
|
def generate(
|
|
self,
|
|
input_ids: torch.Tensor,
|
|
max_new_tokens: int,
|
|
temperature: float = 0.0,
|
|
eos_token_id: int | None = None,
|
|
):
|
|
for _ in range(max_new_tokens):
|
|
logits = self(input_ids)[:, -1]
|
|
if temperature > 0:
|
|
probs = F.softmax(logits / temperature, dim=-1)
|
|
next_token = torch.multinomial(probs, 1)
|
|
else:
|
|
next_token = logits.argmax(-1, keepdim=True)
|
|
input_ids = torch.cat((input_ids, next_token), dim=1)
|
|
if eos_token_id is not None and (next_token.squeeze(-1) == eos_token_id).all():
|
|
break
|
|
return input_ids
|