Files
enlace/tests/test_model.py
T
msaldain 4048936067 Corregir lo que reportó ruff, que nunca se había ejecutado
ruff estaba configurado en pyproject.toml desde el primer commit y jamás se
había corrido. Tenía 15 hallazgos. Dos importan más allá del estilo:

- zip() sin strict= trunca en silencio al más corto. En las comparaciones de
  lotes eso significa que un test podía pasar sin haber comparado todo. Donde
  los largos deben coincidir ahora es strict=True; donde difieren a propósito
  (pares consecutivos) queda strict=False, que documenta la intención.
- Un import sin usar delataba algo peor: BraveBackend se había escrito sin una
  sola prueba. Se agregan ocho, contra una respuesta con la forma que devuelve
  la API, incluidas la limpieza de etiquetas, el caso de límite de tasa —que
  tiene que distinguirse de 'no respondió'— y que la credencial viaje en la
  cabecera y nunca en la URL. El respaldo tiene que funcionar justo cuando el
  primario ya falló; merecía la misma cobertura.

El resto es orden de imports, collections.abc y líneas largas.
2026-07-28 07:25:51 -03:00

154 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 pathlib import Path
from enlace.config.load import load_config
cfg = load_config(Path(__file__).resolve().parents[1] / "configs/runs/pretrain-2060.yaml")
model = Transformer(cfg.model, "math")
assert 45e6 < model.num_parameters() < 55e6