Etapa 0: entorno, configuración validada, modelo y entrenador

Base del proyecto ENLACE: un modelo de lenguaje propio entrenado desde cero,
en español, para asistencia general y familiar. El plan completo está en
docs/PLAN.md.

Esta etapa establece el andamiaje y lo verifica de punta a punta:

- Configuración por capas (hardware × model × train × data) validada con
  pydantic. Ningún hiperparámetro vive en el código y una config inválida
  falla al arrancar, no a las tres horas de entrenamiento.
- Perfiles de hardware que aíslan el salto de GPU: la RTX 2060 (Turing) no
  soporta bfloat16 ni FlashAttention-2, así que entrena en float16 con
  GradScaler y backend mem_efficient; el perfil de la 5090 ya está escrito.
  backends.py valida el perfil contra la GPU real antes de empezar.
- Transformer decoder-only estilo Llama: RMSNorm, SwiGLU, RoPE, GQA,
  embeddings atados, QK-norm y z-loss. Los dos últimos son lo que mantiene
  estable el entrenamiento en float16.
- Entrenador con schedule WSD, acumulación de gradiente, precisión mixta,
  checkpointing atómico y reanudación exacta.
- Cargadores de datos con estado serializable: bytes para el smoke test y
  shards uint16 para el corpus real.

48 tests, entre ellos el crítico: reanudar desde un checkpoint reproduce los
pesos de una corrida ininterrumpida, parámetro por parámetro.

Verificado en CPU: 300 pasos sobre texto en español, loss 3.07 -> 1.63.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
2026-07-27 23:04:45 -03:00
commit 90ac2a6582
33 changed files with 2838 additions and 0 deletions
+152
View File
@@ -0,0 +1,152 @@
"""Propiedades de la arquitectura que tienen que valer en cualquier escala."""
from __future__ import annotations
import math
import pytest
import torch
from enlace.config.schema import ModelConfig
from enlace.model.transformer import Transformer, apply_rope, build_rope_cache
def _cfg(**parches) -> ModelConfig:
base = dict(
name="test",
vocab_size=256,
n_layer=2,
n_head=4,
n_kv_head=2,
d_model=64,
seq_len=32,
)
base.update(parches)
return ModelConfig(**base)
def test_loss_inicial_es_la_del_azar():
"""Un modelo recién inicializado no sabe nada: su loss es ln(vocab).
Si arranca muy por debajo hay una fuga de información (targets sin
desplazar, por ejemplo); si arranca muy por encima, la inicialización está
mal escalada y la corrida va a tardar en despegar.
"""
torch.manual_seed(0)
cfg = _cfg()
model = Transformer(cfg, "math")
x = torch.randint(0, cfg.vocab_size, (8, cfg.seq_len + 1))
_, loss = model(x[:, :-1], x[:, 1:])
assert loss.item() == pytest.approx(math.log(cfg.vocab_size), abs=0.15)
def test_embeddings_atados_comparten_memoria():
model = Transformer(_cfg(tie_embeddings=True), "math")
assert model.lm_head.weight.data_ptr() == model.tok_emb.weight.data_ptr()
suelto = Transformer(_cfg(tie_embeddings=False), "math")
assert suelto.lm_head.weight.data_ptr() != suelto.tok_emb.weight.data_ptr()
def test_atar_embeddings_ahorra_una_tabla_entera():
atado = Transformer(_cfg(tie_embeddings=True), "math").num_parameters()
suelto = Transformer(_cfg(tie_embeddings=False), "math").num_parameters()
assert suelto - atado == 256 * 64
def test_la_atencion_es_causal():
"""Cambiar un token no puede alterar las predicciones anteriores.
Es la propiedad que hace que el entrenamiento tenga sentido; si se rompe,
el loss baja de forma espectacular y el modelo no sirve para nada.
"""
torch.manual_seed(0)
cfg = _cfg()
model = Transformer(cfg, "math").eval()
x = torch.randint(0, cfg.vocab_size, (1, cfg.seq_len))
with torch.no_grad():
base, _ = model(x)
alterado = x.clone()
alterado[0, -1] = (alterado[0, -1] + 1) % cfg.vocab_size
otro, _ = model(alterado)
assert torch.allclose(base[:, :-1], otro[:, :-1], atol=1e-6)
def test_ignore_index_excluye_posiciones_enmascaradas():
"""El SFT entrena solo sobre los turnos del asistente: el resto va con -100."""
torch.manual_seed(0)
cfg = _cfg()
model = Transformer(cfg, "math")
x = torch.randint(0, cfg.vocab_size, (4, cfg.seq_len + 1))
inp, tgt = x[:, :-1], x[:, 1:].clone()
tgt[:, : cfg.seq_len // 2] = -100
_, loss = model(inp, tgt)
assert torch.isfinite(loss)
def test_todo_enmascarado_no_rompe():
cfg = _cfg()
model = Transformer(cfg, "math")
x = torch.randint(0, cfg.vocab_size, (2, cfg.seq_len))
tgt = torch.full_like(x, -100)
_, loss = model(x, tgt)
assert torch.isnan(loss) or torch.isfinite(loss) # no debe explotar
def test_rechaza_secuencias_mas_largas_que_el_contexto():
cfg = _cfg()
model = Transformer(cfg, "math")
x = torch.randint(0, cfg.vocab_size, (1, cfg.seq_len + 1))
with pytest.raises(ValueError, match="supera seq_len"):
model(x)
def test_generate_agrega_exactamente_los_tokens_pedidos():
cfg = _cfg()
model = Transformer(cfg, "math")
inicio = torch.zeros((2, 3), dtype=torch.long)
out = model.generate(inicio, max_new_tokens=7, top_k=5)
assert out.shape == (2, 10)
assert torch.equal(out[:, :3], inicio)
def test_gqa_reduce_las_proyecciones_kv():
"""Con GQA 4:1 las matrices de keys/values son un cuarto de las de queries."""
cfg = _cfg(n_head=8, n_kv_head=2, d_model=64)
model = Transformer(cfg, "math")
attn = model.blocks[0].attn
assert attn.wq.out_features == 8 * cfg.head_dim
assert attn.wk.out_features == 2 * cfg.head_dim
assert attn.n_rep == 4
def test_rope_preserva_la_norma():
"""RoPE es una rotación: no cambia la magnitud de los vectores."""
cos, sin = build_rope_cache(16, 8, 10_000.0)
x = torch.randn(2, 3, 16, 8)
y = apply_rope(x, cos, sin)
assert torch.allclose(x.norm(dim=-1), y.norm(dim=-1), atol=1e-5)
def test_z_loss_penaliza_logits_grandes():
torch.manual_seed(0)
cfg_con = _cfg(z_loss_weight=1.0)
cfg_sin = _cfg(z_loss_weight=0.0)
torch.manual_seed(0)
con = Transformer(cfg_con, "math")
torch.manual_seed(0)
sin = Transformer(cfg_sin, "math")
x = torch.randint(0, 256, (4, 33))
_, loss_con = con(x[:, :-1], x[:, 1:])
_, loss_sin = sin(x[:, :-1], x[:, 1:])
assert loss_con.item() > loss_sin.item()
def test_el_modelo_del_plan_pesa_lo_esperado():
"""tiny-50m tiene que estar cerca de 50M: es lo que entra en la 2060."""
from enlace.config.load import load_config
from pathlib import Path
cfg = load_config(Path(__file__).resolve().parents[1] / "configs/runs/pretrain-2060.yaml")
model = Transformer(cfg.model, "math")
assert 45e6 < model.num_parameters() < 55e6