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>
This commit is contained in:
@@ -0,0 +1,80 @@
|
||||
"""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"
|
||||
Reference in New Issue
Block a user