tiny-agent-112m / code /tiny_agent /checkpoint.py
darioooooo0o's picture
tiny-agent-112m: base + RL weights, tokenizer, code, model card
4397e12 verified
Raw History Blame Contribute Delete
1.71 kB
"""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)