HAKO upload: hako/train/restore.py
Browse files- hako/train/restore.py +95 -0
hako/train/restore.py
ADDED
|
@@ -0,0 +1,95 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Restore a HAKOSystem from a checkpoint (npz + meta) for phase resume.
|
| 2 |
+
|
| 3 |
+
The checkpoint stores every trainable block as numpy arrays (see
|
| 4 |
+
HAKOSystem.state_blocks); restore re-instantiates the components, registers
|
| 5 |
+
the decomposed sources (adapters re-seeded then overwritten by checkpoint
|
| 6 |
+
values), rebuilds the GHSOM grid geometry from the stored W tensors, and
|
| 7 |
+
reloads the scalar state (step counter, growth, game level).
|
| 8 |
+
"""
|
| 9 |
+
from __future__ import annotations
|
| 10 |
+
|
| 11 |
+
import logging
|
| 12 |
+
from pathlib import Path
|
| 13 |
+
|
| 14 |
+
import numpy as np
|
| 15 |
+
import torch
|
| 16 |
+
|
| 17 |
+
from hako.core.ghsom import GHSOMNode
|
| 18 |
+
from hako.core.heads import neuron_key
|
| 19 |
+
from hako.train.trainer import HAKOSystem
|
| 20 |
+
|
| 21 |
+
log = logging.getLogger("hako.restore")
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def restore_system(cfg, tel, mem, ckpt_path: Path) -> tuple:
|
| 25 |
+
sys_ = HAKOSystem(cfg, tel, mem)
|
| 26 |
+
|
| 27 |
+
# sources (register => adapters seeded; checkpoint overwrites below)
|
| 28 |
+
qwen_npz = Path(cfg.artifacts_dir) / "qwen05b_decomposed.npz"
|
| 29 |
+
gran_npz = Path(cfg.artifacts_dir) / "granite1b_decomposed.npz"
|
| 30 |
+
have = {}
|
| 31 |
+
if qwen_npz.exists():
|
| 32 |
+
sys_.register_source("qwen", qwen_npz, cfg.qwen_d)
|
| 33 |
+
have["qwen"] = True
|
| 34 |
+
if gran_npz.exists():
|
| 35 |
+
sys_.register_source("granite", gran_npz, cfg.granite_proto_dim)
|
| 36 |
+
have["granite"] = True
|
| 37 |
+
|
| 38 |
+
from hako.checkpoint import load_state
|
| 39 |
+
state = load_state(ckpt_path)
|
| 40 |
+
block = state.get("moe", {})
|
| 41 |
+
if block:
|
| 42 |
+
sys_.moe.load_numpy({k: v.astype(np.float32) for k, v in block.items()})
|
| 43 |
+
block = state.get("heads", {})
|
| 44 |
+
if block:
|
| 45 |
+
sys_.heads.load_numpy({k: v.astype(np.float32) for k, v in block.items()})
|
| 46 |
+
block = state.get("chain", {})
|
| 47 |
+
if block:
|
| 48 |
+
with torch.no_grad():
|
| 49 |
+
sys_.chain.theta.copy_(torch.as_tensor(block["theta"]))
|
| 50 |
+
sys_.chain.L_hat.copy_(torch.as_tensor(block["L_hat"]))
|
| 51 |
+
sys_.chain.W_beta.copy_(torch.as_tensor(block["W_beta"]))
|
| 52 |
+
block = state.get("tgnn", {})
|
| 53 |
+
if block:
|
| 54 |
+
with torch.no_grad():
|
| 55 |
+
sys_.tgnn.W_att.copy_(torch.as_tensor(block["W_att"]))
|
| 56 |
+
sys_.tgnn.W_msg.copy_(torch.as_tensor(block["W_msg"]))
|
| 57 |
+
block = state.get("router", {})
|
| 58 |
+
if block:
|
| 59 |
+
sd = {k: torch.as_tensor(v, dtype=torch.float32)
|
| 60 |
+
for k, v in block.items()}
|
| 61 |
+
sys_.router.load_state_dict(sd, strict=False)
|
| 62 |
+
block = state.get("diffusion", {})
|
| 63 |
+
if block:
|
| 64 |
+
sd = {k: torch.as_tensor(v, dtype=torch.float32)
|
| 65 |
+
for k, v in block.items()}
|
| 66 |
+
sys_.diffusion.load_state_dict(sd, strict=False)
|
| 67 |
+
gblock = state.get("ghsom", {})
|
| 68 |
+
if gblock:
|
| 69 |
+
for key, arr in gblock.items():
|
| 70 |
+
nid = int(key.split("[")[1].rstrip("]"))
|
| 71 |
+
W = torch.as_tensor(arr, dtype=torch.float32)
|
| 72 |
+
if nid not in sys_.ghsom.nodes:
|
| 73 |
+
node = GHSOMNode(nid, None, None, 0, W)
|
| 74 |
+
sys_.ghsom.nodes[nid] = node
|
| 75 |
+
else:
|
| 76 |
+
sys_.ghsom.nodes[nid].W = W
|
| 77 |
+
sys_.ghsom.nodes[nid].rows, sys_.ghsom.nodes[nid].cols = \
|
| 78 |
+
W.shape[0], W.shape[1]
|
| 79 |
+
sys_.ghsom.root_id = min(sys_.ghsom.nodes)
|
| 80 |
+
sys_.ghsom.next_id = max(sys_.ghsom.nodes) + 1
|
| 81 |
+
block = state.get("theta", {})
|
| 82 |
+
if len(block):
|
| 83 |
+
sys_.theta.theta = np.asarray(block["theta"], dtype=np.float32)
|
| 84 |
+
block = state.get("scalars", {})
|
| 85 |
+
if block:
|
| 86 |
+
sys_.step = int(block.get("step", 0))
|
| 87 |
+
sys_.ghsom.growth_events = int(block.get("growth_events", 0))
|
| 88 |
+
# attach heads for all grid neurons and rebuild optimizer
|
| 89 |
+
for nid, node in sys_.ghsom.nodes.items():
|
| 90 |
+
for idx in range(node.rows * node.cols):
|
| 91 |
+
sys_.heads.attach(neuron_key(nid, idx))
|
| 92 |
+
sys_.rebuild_optimizer()
|
| 93 |
+
log.info("restored system from %s (step=%d, sources=%s)",
|
| 94 |
+
Path(ckpt_path).name, sys_.step, sorted(have))
|
| 95 |
+
return sys_, have
|