
SFT do TatuEngine — Adafactor e os 4 fixes do NaN
Contexto
O pipeline SFT (Supervised Fine-Tuning) do TatuEngine é o gargalo atual do projeto. O Teacher LoRA (Qwen2.5-3B-Instruct + LoRA R=16) já destilou 381 exemplos de raciocínio em 3 datasets. O Student (BitMamba-1B PyTorch) tem o modelo — 1.446B parâmetros, 48 camadas, d_model=2048, d_state=128 — mas não aprende direito.
Não é falta de dados. O problema é mais fundamental: o treinamento explode em NaN sistematicamente.
Foram 4 bugs distintos. Cada um matava o treino de um jeito diferente. E o Adafactor (substituto do AdamW) foi parte da solução e parte do problema ao mesmo tempo.
Os 4 Fixes do NaN
Fix #1 — Softplus overflow no dt_bias: expf(89.0) = inf
O primeiro NaN vinha do SSM step — camada 1, primeiro batch, sempre no ssm_step. O culpado: softplus(dt_bias + dt).
// Antes (explodia):
float dt = log1pf(expf(x)); // softplus(x) = log(1 + exp(x))
// Depois (estável):
float dt = x > 0.0f ? x + log1pf(expf(-x)) : log1pf(expf(x));
// Equivalente numericamente estável:
// softplus(x) = max(0, x) + log(1 + exp(-|x|))
expf(89.0) ≈ 5.5e38 → +inf em float32. Quando dt_bias passava de ~89 após algumas iterações, o softplus retornava +inf, o ssm_step propagava inf, e o backward transformava tudo em NaN.
2 dias de debug. Um patch de 1 linha. log1pf(expf(-|x|)) resolve.
Fix #2 — Adafactor eps=(None, 0.001): o rsqrt(0) em bf16
Adafactor fatora o segundo momento em rank-1 (row + col), economizando ~4 GB vs AdamW (4 GB vs 8.14 GB na RTX 3060). Mas tem um catch:
# Antes (NaN no step 3):
optimizer = torch.optim.Adafactor(
model.parameters(),
lr=3e-5,
eps=(1e-30, 0.001), # eps1 padrão
)
# Depois (estável):
optimizer = torch.optim.Adafactor(
model.parameters(),
lr=3e-5,
eps=(None, 0.001), # eps1 = finfo(bf16).eps ≈ 0.0078
)
O eps1=1e-30 padrão existe pro caso float32 (onde finfo(float32).eps = 1.19e-7). Mas em bf16, finfo(bf16).eps = 0.0078. Quando o segundo momento acumulado é pequeno (primeiros steps), rsqrt(accum + 1e-30) retorna um número enorme porque 1e-30 é subnormal em bf16 e vira 0. Resultado: inf * 0 = NaN.
Passando eps=(None, 0.001), o PyTorch usa finfo(dtype).eps como eps1, que pra bf16 é 0.0078 — o suficiente pra evitar underflow no rsqrt.
Fix #3 — State LR Scale: A_log, D e dt_bias precisam de 50× mais LR
Os parâmetros de estado do SSM (A_log, D, dt_bias) controlam a dinâmica temporal do modelo. Com a LR base (lr=3e-5), eles mal se movem:
# Parâmetros de estado do SSM — precisam de LR maior
STATE_PARAMS = ["A_log", "D", "dt_bias"]
STATE_LR_SCALE = 50 # LR = base × 50
optimizer = torch.optim.Adafactor(
[
{"params": model.core_params()},
{"params": model.state_params(), "lr": args.lr * STATE_LR_SCALE},
],
lr=3e-5,
eps=(None, 0.001),
)
Sem isso, o dt_bias congela perto de 0, o softplus produz log(2) ≈ 0.69 pra sempre, e o modelo não desenvolve discretização temporal — o SSM essencialmente não funciona.
Fix #4 — Gradient Clamping: clamp_(-1, 1) + max-norm
O gradiente do SSM é inerentemente instável por causa da retropropagação através da discretização exp(Δ · A):
# Full SFT: clamp + max-norm (evita overflow bf16)
if not args.no_grad_clamp:
for p in model.parameters():
if p.grad is not None:
p.grad.data.clamp_(-1.0, 1.0) # evita gradientes explosivos
if args.grad_clip > 0:
torch.nn.utils.clip_grad_norm_(
model.parameters(), args.grad_clip, norm_type=float('inf')
)
Em head-only mode (treinando só o lm_head), os gradientes ficam livres — sem clamp porque a head precisa de movimento rápido pra domar a variância da saída. Em full SFT, o clamp [-1, 1] é obrigatório: sem ele, layers profundas explodem no step 5-10.
Tabela Comparativa: Adafactor vs AdamW
| Aspecto | AdamW | Adafactor |
|---|---|---|
| Memória do otimizador | 8.14 GB (2 estados + momentum) | ~4 GB (rank-1 fatorado) |
| Custo por step | Maior (full matrix update) | Menor (row/col update) |
| Estabilidade bf16 | eps=1e-8 seguro |
Precisa de eps=(None, X) tuning |
| Warmup | Necessário (momemtum frio) | Built-in (relative_step=True) |
| Convergência | Mais rápida em transformers | Mais lenta em SSM (precisa state_lr_scale) |
A economia de 4 GB é crítica na RTX 3060 (12 GB). Com AdamW, o modelo 1.4B + optimizer ocupam ~11.5 GB — sem espaço pra batch > 1 ou sequências > 256 tokens. Com Adafactor, cabem 7.5 GB e sobra margem.
O Pipeline Atual
Teacher LoRA (Qwen2.5-3B)
↓ 381 exemplos destilados
Student BitMamba-1B PyTorch
↓ Full warmup (200 steps, plain text)
Checkpoint step_0200
↓ SFT Fase 2 (domain-balanced + replay buffer 9:1)
Student SFT Final
3 datasets: student_train.jsonl (formato prompt + teacher_response), distill_clean_train.jsonl (com filtro de score), tatu_train_data.jsonl (formato híbrido com thought/answer). O preparador de dataset (prepare_datasets()) normaliza os 3 formatos e aplica [THOUGHT]/[ANSWER] tags.
O replay buffer de fineweb-edu (9:1 de raciocínio:pretrain) mantém o modelo de esquecer linguagem geral durante o SFT.
Aprendizados
-
Softplus não é seguro em float32 — a formulação ingênua
log(1 + exp(x))explode quando x > 89. Sempre usar a versão numericamente estável. -
Adafactor economiza VRAM mas cobra em tuning — o eps1 precisa ser calibrado pro dtype.
Noneresolve (delega pro finfo), mas nunca use o default 1e-30 em bf16. -
Parâmetros de estado do SSM são uma classe própria —
dt_bias,A_log,Dcontrolam a discretização temporal e precisam de LR muito maior que os parâmetros de projeção. Ostate_lr_scale=50foi descoberto empiricamente perdendo 3 treinos. -
Gradient clamping não é opcional em SSM — a retropropagação através de
exp(Δ · A)amplifica gradientes exponencialmente.clamp_(-1, 1)+clip_grad_norm_são camadas de segurança complementares, não redundantes.
O SFT completo do Student ainda não rodou (3 épocas, 381 exemplos). Mas agora o pipeline não explode mais — o próximo passo é deixar treinar e ver o que sai.
TL;DR: O Adafactor economizou 4 GB de VRAM (crítico na RTX 3060), mas custou 3 treinos perdidos em NaN até acertar os 4 fixes. Softplus numericamente estável, eps1 calibrado pro bf16, state_lr_scale=50 pros parâmetros do SSM, e clamp+clip como camadas de segurança complementares. O pipeline não explode mais — agora é deixar rodar.