Cccccz's picture
Add files using upload-large-folder tool
d5e0d8f verified
Raw
History Blame Contribute Delete
3.04 kB
"""Atomic Predictor-v4 inference and resumable training checkpoints."""
from __future__ import annotations
import json
import os
import random
from pathlib import Path
from typing import Any
import torch
from safetensors.torch import load_file, save_file
def unwrap_model(model: torch.nn.Module) -> torch.nn.Module:
return model.module if hasattr(model, "module") else model
def trainable_state_dict(
model: torch.nn.Module,
*,
floating_dtype: torch.dtype | None = None,
) -> dict[str, torch.Tensor]:
model = unwrap_model(model)
trainable_names = {
name for name, parameter in model.named_parameters() if parameter.requires_grad
}
return {
name: tensor.detach()
.to(
device="cpu",
dtype=(
floating_dtype
if floating_dtype is not None and tensor.is_floating_point()
else tensor.dtype
),
)
.contiguous()
for name, tensor in model.state_dict().items()
if name in trainable_names
}
def save_predictor_weights(
model: torch.nn.Module,
path: str | Path,
*,
metadata: dict[str, Any],
floating_dtype: torch.dtype = torch.bfloat16,
) -> Path:
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
temporary = path.with_suffix(path.suffix + f".tmp.{os.getpid()}")
save_file(
trainable_state_dict(model, floating_dtype=floating_dtype),
str(temporary),
metadata={
"format": "self_forcing_predictor_v4",
"config": json.dumps(metadata, ensure_ascii=False, sort_keys=True),
},
)
os.replace(temporary, path)
return path
def load_predictor_weights(model: torch.nn.Module, path: str | Path) -> None:
state = load_file(str(path), device="cpu")
result = unwrap_model(model).load_state_dict(state, strict=False)
trainable = {
name
for name, parameter in unwrap_model(model).named_parameters()
if parameter.requires_grad
}
missing_trainable = sorted(trainable.intersection(result.missing_keys))
if result.unexpected_keys or missing_trainable:
raise RuntimeError(
"Predictor weight mismatch: "
f"unexpected={result.unexpected_keys}, "
f"missing_trainable={missing_trainable}"
)
def atomic_torch_save(payload: dict[str, Any], path: str | Path) -> Path:
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
temporary = path.with_suffix(path.suffix + f".tmp.{os.getpid()}")
torch.save(payload, temporary)
os.replace(temporary, path)
return path
def capture_rng_state() -> dict[str, Any]:
return {
"python": random.getstate(),
"torch_cpu": torch.get_rng_state(),
"torch_cuda": torch.cuda.get_rng_state(),
}
def restore_rng_state(state: dict[str, Any]) -> None:
random.setstate(state["python"])
torch.set_rng_state(state["torch_cpu"])
torch.cuda.set_rng_state(state["torch_cuda"])