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"]