"""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