diff --git a/enlace/train/checkpoint.py b/enlace/train/checkpoint.py index e9bd646..f686d9b 100644 --- a/enlace/train/checkpoint.py +++ b/enlace/train/checkpoint.py @@ -36,11 +36,19 @@ def _rng_state() -> dict[str, Any]: def _restore_rng(state: dict[str, Any]) -> None: + """Restaura los generadores aleatorios. + + Los estados se fuerzan a CPU a propósito. `torch.load(map_location="cuda")` + mueve *todos* los tensores del checkpoint a la placa, incluidos estos, y + `set_rng_state` exige un ByteTensor en CPU: sin el `.cpu()` reanudar en GPU + falla con "RNG state must be a torch.ByteTensor". En CPU el error no existe, + así que es un fallo que solo aparece en el servidor. + """ random.setstate(state["python"]) np.random.set_state(state["numpy"]) - torch.set_rng_state(state["torch"]) + torch.set_rng_state(state["torch"].cpu()) if "cuda" in state and torch.cuda.is_available(): - torch.cuda.set_rng_state_all(state["cuda"]) + torch.cuda.set_rng_state_all([s.cpu() for s in state["cuda"]]) def save( diff --git a/tests/test_resume.py b/tests/test_resume.py index 67eab0b..cfbc0b2 100644 --- a/tests/test_resume.py +++ b/tests/test_resume.py @@ -78,3 +78,26 @@ def test_checkpoint_atomico_no_deja_archivos_truncados(tmp_path, texto_es): 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