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>
This commit is contained in:
2026-07-28 01:29:48 -03:00
parent bfa372d9c3
commit 073209d48e
2 changed files with 33 additions and 2 deletions
+23
View File
@@ -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