Los checkpoints se cargan siempre en CPU
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 <noreply@anthropic.com>
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"]
|
||||
|
||||
Reference in New Issue
Block a user