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:
dela
2026-08-25 20:09:27 +08:00
parent 7a12f61de1
commit 49aede9cb2
7 changed files with 117 additions and 36 deletions
+14 -1
View File
@@ -20,6 +20,7 @@ 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")
@@ -247,6 +248,7 @@ class BlockAttnResStack(nn.Module):
if is_final_aggregate
else None
)
self.gradient_checkpointing = False
def forward_naive(self, x: Tensor) -> Tensor:
blocks = [x] # b_0=embedding/input representation
@@ -294,9 +296,20 @@ class BlockAttnResStack(nn.Module):
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)
blocks.append(self._run_block_two_phase(blocks, start, end))
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 (