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:
@@ -36,11 +36,19 @@ def _rng_state() -> dict[str, Any]:
|
|||||||
|
|
||||||
|
|
||||||
def _restore_rng(state: dict[str, Any]) -> None:
|
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"])
|
random.setstate(state["python"])
|
||||||
np.random.set_state(state["numpy"])
|
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():
|
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(
|
def save(
|
||||||
|
|||||||
@@ -78,3 +78,26 @@ def test_checkpoint_atomico_no_deja_archivos_truncados(tmp_path, texto_es):
|
|||||||
run_dir = train(cfg)
|
run_dir = train(cfg)
|
||||||
assert checkpoint.latest(run_dir) is not None
|
assert checkpoint.latest(run_dir) is not None
|
||||||
assert not list(run_dir.glob("*.tmp")), "quedaron temporales de escritura"
|
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
|
||||||
|
|||||||
Reference in New Issue
Block a user