Files
enlace/tests/test_schedules.py
T
msaldain 90ac2a6582 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>
2026-07-27 23:04:45 -03:00

57 lines
1.7 KiB
Python

"""Forma del schedule WSD."""
from __future__ import annotations
from enlace.config.schema import ScheduleConfig
from enlace.train.schedules import build_lr_fn
def _lr_fn(max_steps=1000, warmup=100, decay=200, min_ratio=0.0, base=1e-3):
cfg = ScheduleConfig(
kind="wsd", warmup_steps=warmup, decay_steps=decay, min_lr_ratio=min_ratio
)
return build_lr_fn(cfg, base, max_steps)
def test_warmup_sube_linealmente_desde_arriba_de_cero():
lr = _lr_fn()
assert lr(0) > 0.0 # el primer paso no se desperdicia con lr=0
assert lr(0) < lr(50) < lr(99)
assert lr(99) == 1e-3
def test_la_fase_estable_es_plana():
lr = _lr_fn()
assert lr(100) == lr(500) == lr(799) == 1e-3
def test_el_decay_baja_de_forma_monotona_hasta_cero():
lr = _lr_fn()
valores = [lr(s) for s in range(800, 1000)]
assert all(a >= b for a, b in zip(valores, valores[1:]))
assert valores[0] == 1e-3
assert valores[-1] < 1e-4
def test_min_lr_ratio_pone_un_piso():
lr = _lr_fn(min_ratio=0.1)
assert lr(999) >= 1e-4 * 0.99
def test_sin_decay_el_lr_queda_plano_hasta_el_final():
lr = _lr_fn(decay=0)
assert lr(999) == 1e-3
def test_extender_la_corrida_mueve_el_inicio_del_decay():
"""Documenta la trampa: max_steps define dónde empieza a decaer el LR.
Reanudar una corrida interrumpida exige el mismo max_steps; extenderla es
una decisión distinta, que hay que tomar ramificando desde un checkpoint
anterior al decay.
"""
corto = _lr_fn(max_steps=1000)
largo = _lr_fn(max_steps=2000)
assert corto(850) < corto(700) # ya está decayendo
assert largo(850) == largo(700) # todavía en la fase estable