4048936067
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.
154 lines
5.0 KiB
Python
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
|