03dadc93da
Cambio mecánico, sin efecto en el comportamiento: la suite pasa igual antes y después. Va en un commit propio para no tapar los cambios con sentido. Se agregan además dos flujos de verificación que corren en cada carga al repositorio: - YAML: yamllint para sintaxis y estilo, más la carga de cada config contra su esquema de pydantic. Son cosas distintas — un YAML puede ser sintácticamente perfecto y estar roto igual, con 'run_nombre' en vez de 'run_name'. Ese paso no instala torch: se verificó que la capa de configuración no lo importa, así que corre en segundos en vez de descargar dos gigas y medio de CUDA. - Python: ruff check, ruff format --check y la suite completa con torch de CPU. Las rutas ignoradas de .yamllint.yml van ancladas con barra inicial. Sin anclar, 'data/' y 'runs/' excluían configs/data/ y configs/runs/ — siete archivos, justo los que más importa revisar — y el linter pasaba en verde sin haber mirado nada. Es el mismo defecto que ya había aparecido en .gitignore.
102 lines
4.0 KiB
Python
102 lines
4.0 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"
|
|
|
|
|
|
def test_el_estado_rng_se_restaura_desde_cpu(monkeypatch):
|
|
"""Regresión de un fallo que solo se ve en GPU.
|
|
|
|
`torch.load(map_location="cuda")` mueve todos los tensores del checkpoint a
|
|
la placa, incluido el estado del generador. `torch.set_rng_state` exige un
|
|
ByteTensor en CPU, así que sin forzarlo la reanudación en GPU falla. Acá se
|
|
verifica el contrato que hace falta: lo que se le pasa a torch está en CPU.
|
|
"""
|
|
from enlace.train.checkpoint import _restore_rng, _rng_state
|
|
|
|
recibidos = []
|
|
original = torch.set_rng_state
|
|
monkeypatch.setattr(torch, "set_rng_state", lambda s: (recibidos.append(s), original(s))[1])
|
|
|
|
_restore_rng(_rng_state())
|
|
assert recibidos, "no se llamó a set_rng_state"
|
|
estado = recibidos[0]
|
|
assert estado.device.type == "cpu"
|
|
assert estado.dtype == torch.uint8
|