Download code/tiny_agent/checkpoint.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 1.71 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/tiny_agent/checkpoint.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/tiny_agent/checkpoint.py
-
curl -L -o checkpoint.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/tiny_agent/checkpoint.py
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) | |