From f0ea6202e6c3aa82079b055b7cfad8d12f18a130 Mon Sep 17 00:00:00 2001 From: Mateo Saldain Date: Tue, 28 Jul 2026 01:31:54 -0300 Subject: [PATCH] Los checkpoints se cargan siempre en CPU MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Reanudar en la 2060 fallaba con 'RNG state must be a torch.ByteTensor', y arreglarlo en un lugar lo movía al siguiente: primero el generador de torch, después el del cargador de datos. La causa común es que cargar el checkpoint apuntando a la placa mueve todos los tensores del payload, y los estados de generadores exigen estar en CPU. En vez de seguir parcheando cada consumidor, se ataca el origen: el payload se carga siempre en CPU y desaparece el parámetro map_location. Los pesos no pierden nada, porque load_state_dict copia dentro de los parámetros existentes, que ya están en el dispositivo correcto, y el optimizador reubica su estado solo. El cargador de datos además fuerza CPU por su cuenta, por si un estado llega desde otro lado. Es una familia de fallos invisible desde Gigastar: sin placa, map_location es cpu y los tensores nunca se mueven. Co-Authored-By: Claude Opus 5 --- enlace/data/loaders.py | 5 ++++- enlace/train/checkpoint.py | 17 ++++++++++++++--- enlace/train/train.py | 2 +- 3 files changed, 19 insertions(+), 5 deletions(-) diff --git a/enlace/data/loaders.py b/enlace/data/loaders.py index d9045bf..ab02ab8 100644 --- a/enlace/data/loaders.py +++ b/enlace/data/loaders.py @@ -89,7 +89,10 @@ class ByteStream: def load_state_dict(self, state: dict[str, Any]) -> None: for k, g in self._gens.items(): if k in state: - g.set_state(state[k]) + # `.cpu()` explícito: si el estado viene de un checkpoint que + # alguien cargó apuntando a la placa, `set_state` exige que el + # ByteTensor esté en CPU y falla con un error poco claro. + g.set_state(state[k].cpu()) @staticmethod def decode(ids: list[int]) -> str: diff --git a/enlace/train/checkpoint.py b/enlace/train/checkpoint.py index f686d9b..3236aa6 100644 --- a/enlace/train/checkpoint.py +++ b/enlace/train/checkpoint.py @@ -103,11 +103,22 @@ def load( model: torch.nn.Module, optimizer: torch.optim.Optimizer | None = None, scaler: torch.amp.GradScaler | None = None, - map_location: str | torch.device = "cpu", ) -> dict[str, Any]: - """Restaura el estado y devuelve los metadatos (step, stream, metrics).""" + """Restaura el estado y devuelve los metadatos (step, stream, metrics). + + **El payload se carga siempre en CPU, a propósito.** No es una limitación: + es lo único correcto. Un `map_location` apuntando a la placa mueve *todos* + los tensores del checkpoint, incluidos los estados de los generadores + aleatorios — los de torch y los de los cargadores de datos —, y esos exigen + ByteTensor en CPU. Cargar a GPU rompe la reanudación con un error que en + Gigastar no se puede reproducir porque ahí no hay placa. + + Los pesos no necesitan el atajo: `load_state_dict` copia dentro de los + parámetros existentes, que ya están en el dispositivo correcto, y el + optimizador reubica su estado solo. + """ path = Path(path) - payload = torch.load(path, map_location=map_location, weights_only=False) + payload = torch.load(path, map_location="cpu", weights_only=False) fmt = payload.get("format") if fmt != CHECKPOINT_FORMAT: diff --git a/enlace/train/train.py b/enlace/train/train.py index 1b52cae..975168b 100644 --- a/enlace/train/train.py +++ b/enlace/train/train.py @@ -147,7 +147,7 @@ def train(cfg: Config, resume: bool = False) -> Path: print(f"[enlace] --resume sin checkpoints en {run_dir}: se empieza de cero") else: meta = checkpoint.load( - ckpt_path, model=model, optimizer=optimizer, scaler=scaler, map_location=device + ckpt_path, model=model, optimizer=optimizer, scaler=scaler ) stream.load_state_dict(meta["stream"]) start_step = meta["step"]