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)