Files
enlace/tests/test_resume.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

81 lines
3.2 KiB
Python

"""La prueba más importante del entrenamiento.
Reanudar desde un checkpoint tiene que producir exactamente los mismos pesos que
una corrida ininterrumpida. Si no, el daño es silencioso: las curvas se ven
bien, el modelo entrena de más sobre datos ya vistos, y nadie se entera hasta
que la corrida de dos días terminó y el resultado no sirve.
"""
from __future__ import annotations
import shutil
from pathlib import Path
import torch
from enlace.config.load import load_config
from enlace.train import checkpoint
from enlace.train.train import train
from tests.conftest import write_run_config
def _final_weights(run_dir, step: int):
payload = torch.load(run_dir / f"ckpt-{step:08d}.pt", map_location="cpu", weights_only=False)
return payload["model"], payload
def test_reanudar_reproduce_la_corrida_completa(tmp_path, texto_es):
# Corrida A: 20 pasos sin interrupción.
cfg_a = load_config(
write_run_config(tmp_path, texto_es, run_name="A", max_steps=20, checkpoint_every=10)
)
dir_a = train(cfg_a)
# Corrida B: la misma config, simulando que el proceso murió en el paso 10.
# Se copia el checkpoint del paso 10 y se reanuda desde ahí.
#
# La config tiene que ser idéntica, incluido max_steps: el schedule WSD
# ubica el inicio del decay en max_steps - decay_steps, así que reanudar
# con otro max_steps cambia el LR de los pasos que faltan. No es un defecto
# del checkpointing, pero sí una trampa fácil de pisar.
cfg_b = load_config(
write_run_config(tmp_path, texto_es, run_name="B", max_steps=20, checkpoint_every=10)
)
dir_b = Path(cfg_b.train.out_dir) / cfg_b.train.run_name
dir_b.mkdir(parents=True, exist_ok=True)
shutil.copy(dir_a / "ckpt-00000010.pt", dir_b / "ckpt-00000010.pt")
train(cfg_b, resume=True)
pesos_a, payload_a = _final_weights(dir_a, 20)
pesos_b, payload_b = _final_weights(dir_b, 20)
assert payload_a["step"] == payload_b["step"] == 20
assert set(pesos_a) == set(pesos_b)
for nombre, tensor_a in pesos_a.items():
assert torch.equal(tensor_a, pesos_b[nombre]), (
f"el parámetro {nombre} difiere tras reanudar: la reanudación no es exacta"
)
def test_el_checkpoint_guarda_la_posicion_del_stream(tmp_path, texto_es):
cfg = load_config(
write_run_config(tmp_path, texto_es, run_name="C", max_steps=4, checkpoint_every=4)
)
run_dir = train(cfg)
payload = torch.load(run_dir / "ckpt-00000004.pt", map_location="cpu", weights_only=False)
# Sin el estado del stream, reanudar volvería al principio del corpus.
assert "stream" in payload and payload["stream"], "el checkpoint no guardó el stream"
assert "train" in payload["stream"]
# Y sin la config, el checkpoint no sería reproducible.
assert payload["config"]["model"]["name"] == "char-smoke"
def test_checkpoint_atomico_no_deja_archivos_truncados(tmp_path, texto_es):
cfg = load_config(
write_run_config(tmp_path, texto_es, run_name="D", max_steps=6, checkpoint_every=6)
)
run_dir = train(cfg)
assert checkpoint.latest(run_dir) is not None
assert not list(run_dir.glob("*.tmp")), "quedaron temporales de escritura"