PowerMachine commited on
Commit
cb7cb4c
·
verified ·
1 Parent(s): 0df4ef1

HAKO upload: hako/train/phase3_nlp.py

Browse files
Files changed (1) hide show
  1. hako/train/phase3_nlp.py +179 -0
hako/train/phase3_nlp.py ADDED
@@ -0,0 +1,179 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Phase 3 -- NLP/NLG on the real multilingual corpus.
2
+
3
+ Steps:
4
+ 1. Train HAKO's OWN byte-BPE (Lemma 1 exact parallel counts) on the
5
+ mixed EN+PT+zh corpus; verify sharded == sequential merge histories.
6
+ 2. Train the tiny causal NLG decoder, conditioned on HAKO fused latents
7
+ (z_cond from the system forward of each prompt), on
8
+ ultra-alpaca-ptbr instruction pairs + doc continuation.
9
+ 3. Sample generations for the audit trail (telemetry/thinking.jsonl).
10
+
11
+ Theorem anchors asserted at runtime:
12
+ Lemma 1 (BPE additivity): merge-history equality sharded vs full.
13
+ T-LM: loss is finite, decreasing EMA on the validation slice.
14
+ """
15
+ from __future__ import annotations
16
+
17
+ import logging
18
+ import time
19
+ from pathlib import Path
20
+
21
+ import numpy as np
22
+ import torch
23
+
24
+ from hako.checkpoint import save_state
25
+ from hako.memory_manager import MemoryManager
26
+ from hako.nlp.datasets import Corpus
27
+ from hako.nlp.nlg_head import NLGHead
28
+ from hako.telemetry import Telemetry
29
+ from hako.tokenizer.byte_bpe import ByteBPETokenizer
30
+ from hako.train.trainer import HAKOSystem
31
+
32
+ log = logging.getLogger("hako.phase3")
33
+
34
+
35
+ def _verify_bpe_parallelism(corpus_docs):
36
+ """Lemma 1 check: small-scale sharded trainer == single-shard trainer."""
37
+ small = [d[:120] for d in corpus_docs[:60] if d]
38
+ tok_full = ByteBPETokenizer(vocab_size=300, min_freq=2, shards=1)
39
+ tok_full.train(small, log_every=0)
40
+ tok_shard = ByteBPETokenizer(vocab_size=300, min_freq=2, shards=4)
41
+ tok_shard.train(small, log_every=0)
42
+ ok = tok_full.merge_history == tok_shard.merge_history
43
+ return ok, len(tok_full.merge_history)
44
+
45
+
46
+ def run(cfg, tel: Telemetry, mem: MemoryManager, deadline_s: int,
47
+ carry: dict) -> dict:
48
+ t0 = time.time()
49
+ sys_: HAKOSystem = carry["system"]
50
+ corpus: Corpus = carry.get("corpus") or Corpus()
51
+ for prompt, resp in corpus.docs.get("pt_instr", []):
52
+ corpus.add_instruction(prompt, resp)
53
+
54
+ # ---------------- 1. byte-BPE (cached if already trained) -------------
55
+ bpe_cache = Path(cfg.artifacts_dir) / "hako_byte_bpe.json"
56
+ if bpe_cache.exists():
57
+ bpe = ByteBPETokenizer.load(bpe_cache)
58
+ tel.orchestration(event="bpe_loaded_from_cache", vocab=bpe.vocab_size)
59
+ else:
60
+ all_docs = sum(corpus.docs.values(), [])
61
+ ok, n_merges = _verify_bpe_parallelism(all_docs)
62
+ assert ok, "Lemma 1 violated: sharded BPE != sequential BPE"
63
+ tel.orchestration(event="bpe_lemma1_verified", merges=int(n_merges))
64
+ bpe = ByteBPETokenizer(vocab_size=cfg.bpe_vocab,
65
+ min_freq=cfg.bpe_min_freq,
66
+ shards=cfg.bpe_shards)
67
+ bpe.train(all_docs, log_every=256)
68
+ bpe.save(bpe_cache)
69
+ tel.orchestration(event="bpe_trained", vocab=bpe.vocab_size)
70
+
71
+ # ---------------- 2. NLG training data --------------------------------
72
+ # instruction pairs from ultra-alpaca-ptbr (prompt, response)
73
+ instr = list(corpus.instructions)
74
+ if not instr and corpus.docs.get("pt"):
75
+ al = corpus.docs["pt"]
76
+ instr = [(d[:200], d[200:600]) for d in al if len(d) > 500]
77
+ train_pairs = instr[: int(len(instr) * 0.92)] or instr
78
+ val_pairs = instr[int(len(instr) * 0.92):] or instr[-64:]
79
+
80
+ nlg = NLGHead(vocab=bpe.vocab_size, N_dim=cfg.N_dim, d=cfg.N_dim,
81
+ layers=cfg.nlg_layers, heads=cfg.nlg_heads,
82
+ seq_len=cfg.seq_len, dropout=cfg.nlg_dropout,
83
+ seed=cfg.seed)
84
+ opt = torch.optim.AdamW(nlg.parameters(), lr=2.5e-4, weight_decay=0.01)
85
+ gen = torch.Generator().manual_seed(cfg.seed + 3)
86
+
87
+ # precompute HAKO latents for prompts ONCE (conditioning cache)
88
+ Z_t = carry["Z_t"]
89
+ _cond_cache: dict = {}
90
+
91
+ def cond_for(i: int) -> torch.Tensor:
92
+ j = i % len(Z_t)
93
+ if j not in _cond_cache:
94
+ with torch.no_grad():
95
+ _cond_cache[j] = sys_.forward_sample(
96
+ {"qwen": Z_t[j]}, train=False)["z_final"]
97
+ return _cond_cache[j]
98
+
99
+ def encode_batch(pairs, start, bs):
100
+ batch_ids, batch_z = [], []
101
+ for k in range(start, min(start + bs, len(pairs))):
102
+ p, r = pairs[k]
103
+ ids = bpe.encode(("P: " + p + "\nR: " + r)[:700])[: cfg.seq_len]
104
+ if len(ids) < 8:
105
+ continue
106
+ batch_ids.append(torch.tensor(ids, dtype=torch.long))
107
+ batch_z.append(cond_for(k))
108
+ return batch_ids, batch_z
109
+
110
+ step = 0
111
+ losses_ema = None
112
+ bs = 12
113
+ cursor = 0
114
+ nlg_t0 = time.time() # NLG budget starts AFTER BPE, not before
115
+ while time.time() - nlg_t0 < deadline_s:
116
+ batch_ids, batch_z = encode_batch(train_pairs, cursor, bs)
117
+ cursor += bs
118
+ if cursor >= len(train_pairs):
119
+ cursor = 0
120
+ if not batch_ids:
121
+ continue
122
+ L = max(len(x) for x in batch_ids)
123
+ pad = torch.zeros(len(batch_ids), L, dtype=torch.long)
124
+ for i, x in enumerate(batch_ids):
125
+ pad[i, : len(x)] = x
126
+ zc = torch.stack(batch_z)
127
+ loss = nlg.loss(pad, zc)
128
+ opt.zero_grad()
129
+ loss.backward()
130
+ torch.nn.utils.clip_grad_norm_(nlg.parameters(), 1.5)
131
+ opt.step()
132
+ step += 1
133
+ losses_ema = float(loss) if losses_ema is None else \
134
+ 0.98 * losses_ema + 0.02 * float(loss)
135
+ if step % 20 == 0:
136
+ tel.learning(step=step, phase="phase3", loss_total=float(loss),
137
+ loss_lm_ema=losses_ema, batch=L)
138
+ mem.check()
139
+ if step % 300 == 0:
140
+ log.info("phase3 nlg step %d ema=%.3f", step, losses_ema)
141
+
142
+ # ---------------- 3. validation + samples ------------------------------
143
+ nlg.eval()
144
+ val_losses = []
145
+ with torch.no_grad():
146
+ for start in range(0, min(len(val_pairs), 96), bs):
147
+ batch_ids, batch_z = encode_batch(val_pairs, start, bs)
148
+ if not batch_ids:
149
+ continue
150
+ L = max(len(x) for x in batch_ids)
151
+ pad = torch.zeros(len(batch_ids), L, dtype=torch.long)
152
+ for i, x in enumerate(batch_ids):
153
+ pad[i, : len(x)] = x
154
+ val_losses.append(float(nlg.loss(pad, torch.stack(batch_z))))
155
+ val_loss = float(np.mean(val_losses)) if val_losses else float("nan")
156
+ samples = []
157
+ demo_prompt = "P: Explique o que e um modelo de Kohonen.\nR:"
158
+ ids0 = bpe.encode(demo_prompt)
159
+ zc0 = cond_for(0).view(1, -1)
160
+ out_ids = nlg.generate(ids0, zc0, max_new=48, temperature=0.95)
161
+ sample_text = bpe.decode(out_ids)
162
+ samples.append(sample_text)
163
+ tel.thinking(request_id="nlg-demo", cycle=1, step=4, complete=True,
164
+ sample=sample_text[:500])
165
+
166
+ # ---------------- 4. checkpoint ----------------------------------------
167
+ path = save_state(Path(cfg.ckpt_dir) / "hako_phase3.npz",
168
+ **sys_.state_blocks())
169
+ np.savez_compressed(Path(cfg.artifacts_dir) / "nlg_head.npz",
170
+ **{k: v.detach().numpy()
171
+ for k, v in nlg.state_dict().items()})
172
+ carry["nlg"] = nlg
173
+ carry["bpe"] = bpe
174
+ carry["nlg_val_loss"] = val_loss
175
+ carry["nlg_sample"] = sample_text
176
+ carry["checkpoint"] = str(path)
177
+ tel.orchestration(event="phase3_done", nlg_steps=int(step),
178
+ nlg_val_loss=val_loss)
179
+ return carry