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