"""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