Files
enlace/tests/test_resume.py
T
msaldain 073209d48e Arreglar la reanudación en GPU: el estado del generador debe volver en CPU
Reanudar una corrida en la 2060 fallaba con 'RNG state must be a
torch.ByteTensor'. La causa: al cargar el checkpoint con map_location apuntando
a la placa, torch.load mueve todos los tensores del payload a la GPU, incluido
el estado de los generadores aleatorios, y set_rng_state exige un ByteTensor en
CPU.

Es un fallo que no se puede ver desde Gigastar: en CPU el map_location es cpu y
los tensores nunca se mueven. Lo encontró la verificación de reanudación exacta
corrida en el servidor, que es justamente el criterio de aceptación del plan.

Se agrega una prueba de regresión que verifica el contrato —lo que recibe torch
está en CPU y es uint8— y que sí corre sin GPU.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-07-28 01:29:48 -03:00

104 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