90ac2a6582
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>
153 lines
5.0 KiB
Python
153 lines
5.0 KiB
Python
"""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
|