File size: 1,709 Bytes
4397e12 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 | """Load snapshots written by scripts/train.py (snap_*.pt / final.pt), full checkpoints, or an exported
directory (model.safetensors + config.json, as published on Hugging Face)."""
from __future__ import annotations
import json
import os
import torch
from tiny_agent.model import ModelConfig, TinyAgentLM
def load_model(path: str, device="xpu", dtype=torch.float32) -> TinyAgentLM:
if os.path.isdir(path):
from safetensors.torch import load_file
st = {"model": load_file(os.path.join(path, "model.safetensors")),
"config": json.load(open(os.path.join(path, "config.json")))}
else:
st = torch.load(path, map_location="cpu", weights_only=False)
cfg_d = st.get("config") or st["meta"]["config"]
for k in ("engram_layers", "engram_orders"):
cfg_d[k] = tuple(cfg_d[k])
model = TinyAgentLM(ModelConfig(**cfg_d))
model.load_state_dict({k: v.float() if v.is_floating_point() else v for k, v in st["model"].items()})
return model.to(device=device, dtype=dtype)
def export(path: str, out_dir: str) -> None:
"""Write a checkpoint as out_dir/model.safetensors (bf16) + out_dir/config.json."""
from safetensors.torch import save_file
st = torch.load(path, map_location="cpu", weights_only=False)
cfg_d = dict(st.get("config") or st["meta"]["config"])
sd = {k: (v.to(torch.bfloat16) if v.is_floating_point() else v).contiguous() for k, v in st["model"].items()}
os.makedirs(out_dir, exist_ok=True)
save_file(sd, os.path.join(out_dir, "model.safetensors"))
json.dump({k: list(v) if isinstance(v, tuple) else v for k, v in cfg_d.items()},
open(os.path.join(out_dir, "config.json"), "w"), indent=1)
|