PowerMachine commited on
Commit
b9ea98c
·
verified ·
1 Parent(s): 96a1aba

HAKO upload: hako/train/restore.py

Browse files
Files changed (1) hide show
  1. 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