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.
50 lines
1.4 KiB
Python
50 lines
1.4 KiB
Python
"""LR scale and token-horizon helpers for train_k3 / train_sft."""
|
|
from __future__ import annotations
|
|
|
|
import math
|
|
|
|
|
|
def lr_scale(
|
|
opt_step: int,
|
|
warmup: int,
|
|
total_opt: int,
|
|
min_ratio: float = 0.1,
|
|
) -> float:
|
|
"""Linear warmup (optimizer steps) then cosine down to ``min_ratio``.
|
|
|
|
``opt_step`` is 0-indexed at the optimizer update that is about to run.
|
|
"""
|
|
if warmup > 0 and opt_step < warmup:
|
|
return (opt_step + 1) / warmup
|
|
denom = max(total_opt - warmup - 1, 1)
|
|
progress = min(max(opt_step - warmup, 0) / denom, 1.0)
|
|
cosine = 0.5 * (1.0 + math.cos(math.pi * progress))
|
|
return min_ratio + (1.0 - min_ratio) * cosine
|
|
|
|
|
|
def tokens_per_micro(batch: int, seq_len: int) -> int:
|
|
return batch * seq_len
|
|
|
|
|
|
def total_opt_steps(
|
|
*,
|
|
max_tokens: int | None,
|
|
max_micro: int | None,
|
|
batch: int,
|
|
seq_len: int,
|
|
grad_acc: int,
|
|
) -> int:
|
|
"""Optimizer-step horizon used by cosine. At least 1.
|
|
|
|
``max_tokens`` is the training budget when set; ``max_micro`` is only used
|
|
when ``max_tokens`` is None. Otherwise a default ``--steps 2000`` would
|
|
shrink a 1B-token cosine to 250 opt steps.
|
|
"""
|
|
acc = max(grad_acc, 1)
|
|
if max_tokens is not None and max_tokens > 0:
|
|
tpm = max(tokens_per_micro(batch, seq_len), 1)
|
|
return max(math.ceil(max_tokens / (tpm * acc)), 1)
|
|
if max_micro is not None and max_micro > 0:
|
|
return max(math.ceil(max_micro / acc), 1)
|
|
return 1
|