From 073209d48ee4142d339d2ed08d32ca5669e9ddae Mon Sep 17 00:00:00 2001 From: Mateo Saldain Date: Tue, 28 Jul 2026 01:29:48 -0300 Subject: [PATCH] =?UTF-8?q?Arreglar=20la=20reanudaci=C3=B3n=20en=20GPU:=20?= =?UTF-8?q?el=20estado=20del=20generador=20debe=20volver=20en=20CPU?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 --- enlace/train/checkpoint.py | 12 ++++++++++-- tests/test_resume.py | 23 +++++++++++++++++++++++ 2 files changed, 33 insertions(+), 2 deletions(-) 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