SFT do TatuEngine — Adafactor e os 4 fixes do NaN
tatuengine·

SFT do TatuEngine — Adafactor e os 4 fixes do NaN

📖 8 min de leitura← Voltar para timeline

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

  1. 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.

  2. Adafactor economiza VRAM mas cobra em tuning — o eps1 precisa ser calibrado pro dtype. None resolve (delega pro finfo), mas nunca use o default 1e-30 em bf16.

  3. Parâmetros de estado do SSM são uma classe própriadt_bias, A_log, D controlam a discretização temporal e precisam de LR muito maior que os parâmetros de projeção. O state_lr_scale=50 foi descoberto empiricamente perdendo 3 treinos.

  4. 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.


~/lifelog — bash
$cat about.txt
╔══════════════════════════════════════╗
║  Samuel Medeiros                    ║
║  Senior Software Engineer           ║
║  Stack: Python · TypeScript · Rust  ║
║  Projetos: Arachne, Dogwalk,        ║
║            Capivara, TatuEngine      ║
╚══════════════════════════════════════╝
      
$