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.
This commit is contained in:
+27
-8
@@ -26,6 +26,28 @@ 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":
|
||||
@@ -89,17 +111,14 @@ class CausalLM(nn.Module):
|
||||
x = activation_checkpoint(block, x, use_reentrant=False)
|
||||
else:
|
||||
x = block(x)
|
||||
elif self.gradient_checkpointing and self.training:
|
||||
x = activation_checkpoint(self.mixer, x, use_reentrant=False)
|
||||
else:
|
||||
self.mixer.gradient_checkpointing = self.gradient_checkpointing
|
||||
x = self.mixer(x)
|
||||
logits = self.lm_head(self.norm(x))
|
||||
hidden = self.norm(x)
|
||||
if labels is None:
|
||||
return logits
|
||||
return F.cross_entropy(
|
||||
logits[:, :-1].reshape(-1, logits.size(-1)),
|
||||
labels[:, 1:].reshape(-1),
|
||||
ignore_index=ignore_index,
|
||||
return self.lm_head(hidden)
|
||||
return _chunked_linear_cross_entropy(
|
||||
hidden, self.lm_head.weight, labels, ignore_index=ignore_index
|
||||
)
|
||||
|
||||
@torch.inference_mode()
|
||||
|
||||
Reference in New Issue
Block a user