Delon Swartz commited on
Commit
ba850eb
·
verified ·
1 Parent(s): 4d1e32e

add kolm.py

Browse files
Files changed (1) hide show
  1. kolm.py +362 -0
kolm.py ADDED
@@ -0,0 +1,362 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ KOLM — Kuramoto Oscillator Language Model (OLM v4).
4
+
5
+ v3 (olm.py) trains damped *amplitude* oscillators (coRNN). v4 swaps the core
6
+ for Kuramoto oscillators, AKOrN-style (Miyato et al., ICLR 2025): each hidden
7
+ unit is a unit vector x_i on the sphere S^{N-1}. Units interact through one
8
+ trained scalar coupling J_ij per pair and synchronize; binding-by-synchrony
9
+ means units that align are bound into one active concept, and the same units
10
+ re-align differently for another concept.
11
+
12
+ One char step (all trained: J, A, U, Wo, bo):
13
+
14
+ y_i = sum_j J_ij x_j + U[char]_i coupling + input drive
15
+ P_i = y_i - (x_i . y_i) x_i project onto tangent space
16
+ x_i <- normalize(x_i + dt (A_i x_i + P_i)) A_i antisymmetric: rotation
17
+
18
+ Unit norm is enforced by construction, so stability is geometric — no
19
+ gamma/eps clipping. Same zero-pretraining setup as v3: learns one novel as
20
+ '?keyword;sentence' lines; chat picks a keyword and generates conditioned
21
+ on it. Pure NumPy, hand-derived BPTT, ~141k params at H=256 N=4.
22
+
23
+ python3 kolm.py --gradcheck
24
+ python3 kolm.py --train --save kolm256.npz
25
+ python3 kolm.py --load kolm256.npz
26
+ """
27
+
28
+ import argparse
29
+ import sys
30
+ import time
31
+ import numpy as np
32
+
33
+ from olm import CHARS, C2I, V, F32, build_corpus, pick_keyword
34
+
35
+
36
+ # ---------------------------------------------------------------- model
37
+ def init_params(H: int, N: int, seed: int = 0, om_init: float = 0.1,
38
+ groups: int = 0):
39
+ rng = np.random.default_rng(seed)
40
+ p = {
41
+ "J": rng.normal(0, 1.0 / np.sqrt(H), (H, H)).astype(F32),
42
+ "Om": rng.normal(0, om_init, (H, N, N)).astype(F32), # A = Om - Om^T
43
+ "U": rng.normal(0, 0.5, (V, H, N)).astype(F32),
44
+ "Wo": rng.normal(0, 1.0 / np.sqrt(H * N),
45
+ (H * N + groups, V)).astype(F32),
46
+ "bo": np.zeros(V, F32),
47
+ }
48
+ if groups:
49
+ # coherence readout: r_g = ||sum_i M_gi x_i|| — a trained order
50
+ # parameter, large only when the units row g weights are in sync
51
+ p["M"] = rng.normal(0, 1.0 / np.sqrt(H), (groups, H)).astype(F32)
52
+ return p
53
+
54
+
55
+ def init_state(B: int, H: int, N: int, dtype=F32):
56
+ x = np.zeros((B, H, N), dtype)
57
+ x[..., 0] = 1.0 # all units start at the same pole; A and U disperse them
58
+ return x
59
+
60
+
61
+ def antisym(Om):
62
+ return Om - Om.transpose(0, 2, 1)
63
+
64
+
65
+ def couple(J, x):
66
+ """y[b,i,:] = sum_j J[i,j] x[b,j,:] as one BLAS matmul."""
67
+ B, H, N = x.shape
68
+ return (J @ x.transpose(1, 0, 2).reshape(H, B * N)) \
69
+ .reshape(-1, B, N).transpose(1, 0, 2)
70
+
71
+
72
+ def readout_feats(p, x):
73
+ """[flat phases, coherence r] and the population vectors m (or None)."""
74
+ B, H, N = x.shape
75
+ if "M" not in p:
76
+ return x.reshape(B, H * N), None, None
77
+ m = couple(p["M"], x) # (B,G,N)
78
+ r = np.sqrt((m * m).sum(-1) + 1e-8) # (B,G)
79
+ return np.concatenate([x.reshape(B, H * N), r], axis=1), m, r
80
+
81
+
82
+ def forward(p, x, X, Y, dt):
83
+ """One BPTT window. X,Y: (B,T) int ids. Returns loss, caches, final state."""
84
+ B, T = X.shape
85
+ H, N = p["Om"].shape[0], p["Om"].shape[1]
86
+ A = antisym(p["Om"])
87
+ caches, loss = [], 0.0
88
+ for t in range(T):
89
+ x0 = x
90
+ y = couple(p["J"], x0) + p["U"][X[:, t]]
91
+ d = (x0 * y).sum(-1, keepdims=True)
92
+ P = y - d * x0
93
+ R = np.einsum("inm,bim->bin", A, x0)
94
+ xt = x0 + dt * (R + P)
95
+ nrm = np.sqrt((xt * xt).sum(-1, keepdims=True))
96
+ x = xt / nrm
97
+ f, m, r = readout_feats(p, x)
98
+ logits = f @ p["Wo"] + p["bo"]
99
+ logits -= logits.max(axis=1, keepdims=True)
100
+ e = np.exp(logits)
101
+ probs = e / e.sum(axis=1, keepdims=True)
102
+ loss -= np.log(probs[np.arange(B), Y[:, t]] + 1e-9).sum()
103
+ caches.append((x0, y, d, nrm, x, f, m, r, probs, X[:, t], Y[:, t]))
104
+ return loss / (B * T), caches, x
105
+
106
+
107
+ def backward(p, caches, dt):
108
+ B = caches[0][0].shape[0]
109
+ T = len(caches)
110
+ H, N = p["Om"].shape[0], p["Om"].shape[1]
111
+ A = antisym(p["Om"])
112
+ g = {k: np.zeros_like(v) for k, v in p.items()}
113
+ gA = np.zeros_like(A)
114
+ gx = np.zeros_like(caches[0][0])
115
+ for t in reversed(range(T)):
116
+ x0, y, d, nrm, x1, f, m, r, probs, xs, ys = caches[t]
117
+ # output head
118
+ dlog = probs.copy()
119
+ dlog[np.arange(B), ys] -= 1.0
120
+ dlog /= (B * T)
121
+ g["Wo"] += f.T @ dlog
122
+ g["bo"] += dlog.sum(0)
123
+ gf = dlog @ p["Wo"].T
124
+ gx1 = gx + gf[:, :H * N].reshape(B, H, N)
125
+ if m is not None:
126
+ # r = ||m||, m_g = sum_i M_gi x1_i
127
+ gm = (gf[:, H * N:] / r)[:, :, None] * m
128
+ gm_r = gm.transpose(1, 0, 2).reshape(-1, B * N)
129
+ x1_r = x1.transpose(1, 0, 2).reshape(H, B * N)
130
+ g["M"] += gm_r @ x1_r.T
131
+ gx1 += (p["M"].T @ gm_r).reshape(H, B, N).transpose(1, 0, 2)
132
+ # x1 = xt / ||xt|| => gxt = (I - x1 x1^T) gx1 / ||xt||
133
+ gxt = (gx1 - (gx1 * x1).sum(-1, keepdims=True) * x1) / nrm
134
+ # xt = x0 + dt*(R + P)
135
+ gx0 = gxt.copy()
136
+ gR = dt * gxt
137
+ gP = dt * gxt
138
+ # R_i = A_i x0_i
139
+ gA += np.einsum("bin,bim->inm", gR, x0)
140
+ gx0 += np.einsum("inm,bin->bim", A, gR)
141
+ # P = y - (x0.y) x0
142
+ gPdot = (gP * x0).sum(-1, keepdims=True)
143
+ gy = gP - gPdot * x0
144
+ gx0 -= gPdot * y + d * gP
145
+ # y = J x0 + U[xs]
146
+ gy_r = gy.transpose(1, 0, 2).reshape(H, B * N)
147
+ x0_r = x0.transpose(1, 0, 2).reshape(H, B * N)
148
+ g["J"] += gy_r @ x0_r.T
149
+ gx0 += (p["J"].T @ gy_r).reshape(H, B, N).transpose(1, 0, 2)
150
+ np.add.at(g["U"], xs, gy)
151
+ gx = gx0
152
+ g["Om"] = gA - gA.transpose(0, 2, 1)
153
+ return g
154
+
155
+
156
+ def gradcheck():
157
+ """Finite-difference check of the hand-derived BPTT gradients."""
158
+ np.random.seed(1)
159
+ H, N, B, T, dt = 5, 3, 2, 4, 0.5
160
+ p = {k: v.astype(np.float64)
161
+ for k, v in init_params(H, N, seed=3, groups=3).items()}
162
+ X = np.random.randint(0, V, (B, T))
163
+ Y = np.random.randint(0, V, (B, T))
164
+ x = np.random.randn(B, H, N)
165
+ x /= np.sqrt((x * x).sum(-1, keepdims=True))
166
+ loss, caches, _ = forward(p, x, X, Y, dt)
167
+ g = backward(p, caches, dt)
168
+ worst = 0.0
169
+ for name in p:
170
+ flat = p[name].reshape(-1)
171
+ for idx in np.random.choice(flat.size, min(6, flat.size), replace=False):
172
+ eps_ = 1e-5
173
+ old = flat[idx]
174
+ flat[idx] = old + eps_
175
+ lp, _, _ = forward(p, x, X, Y, dt)
176
+ flat[idx] = old - eps_
177
+ lm, _, _ = forward(p, x, X, Y, dt)
178
+ flat[idx] = old
179
+ num = (lp - lm) / (2 * eps_)
180
+ ana = g[name].reshape(-1)[idx]
181
+ rel = abs(num - ana) / max(1e-8, abs(num) + abs(ana))
182
+ worst = max(worst, rel)
183
+ print(f"gradcheck worst relative error: {worst:.2e} "
184
+ f"({'PASS' if worst < 1e-4 else 'FAIL'})")
185
+ return worst < 1e-4
186
+
187
+
188
+ # ---------------------------------------------------------------- training
189
+ def adam_step(p, g, m, v, t, lr, clip=1.0):
190
+ norm = np.sqrt(sum(float((gi ** 2).sum()) for gi in g.values()))
191
+ scale = min(1.0, clip / (norm + 1e-8))
192
+ b1, b2, e = 0.9, 0.999, 1e-8
193
+ for k in p:
194
+ gk = g[k] * scale
195
+ m[k] = b1 * m[k] + (1 - b1) * gk
196
+ v[k] = b2 * v[k] + (1 - b2) * gk * gk
197
+ mh = m[k] / (1 - b1 ** t)
198
+ vh = v[k] / (1 - b2 ** t)
199
+ p[k] -= (lr * mh / (np.sqrt(vh) + e)).astype(F32)
200
+ return norm
201
+
202
+
203
+ def train(args):
204
+ corpus, freq, kwc = build_corpus(args.text)
205
+ ids = np.array([C2I[c] for c in corpus], dtype=np.int64)
206
+ n_val = 20_000
207
+ tr, va = ids[:-n_val], ids[-n_val:]
208
+ B, T, H, N, dt = args.batch, args.seq, args.hidden, args.ndim, args.dt
209
+
210
+ p = init_params(H, N, om_init=args.om_init, groups=args.groups)
211
+ n_params = sum(x.size for x in p.values())
212
+ print(f"KOLM v4 | hidden {H} x S^{N - 1} | params {n_params:,} "
213
+ f"({n_params * 4 / 1e6:.2f} MB) | corpus {len(ids):,} chars", flush=True)
214
+
215
+ L = len(tr) // B
216
+ streams = tr[: B * L].reshape(B, L)
217
+ m = {k: np.zeros_like(x) for k, x in p.items()}
218
+ v = {k: np.zeros_like(x) for k, x in p.items()}
219
+ step = 0
220
+ for ep in range(1, args.epochs + 1):
221
+ lr = args.lr * (0.5 ** max(0, ep - args.epochs + 6) if ep > args.epochs - 6 else 1.0)
222
+ x = init_state(B, H, N)
223
+ tot = nb = 0
224
+ t0 = time.time()
225
+ for s in range(0, L - T - 1, T):
226
+ X, Y = streams[:, s:s + T], streams[:, s + 1:s + T + 1]
227
+ loss, caches, x = forward(p, x, X, Y, dt)
228
+ g = backward(p, caches, dt)
229
+ step += 1
230
+ adam_step(p, g, m, v, step, lr)
231
+ tot += loss
232
+ nb += 1
233
+ vl, vacc = evaluate(p, va, dt)
234
+ print(f"epoch {ep:2d} | train loss {tot / nb:.3f} | "
235
+ f"val loss {vl:.3f} | val acc {vacc:.1%} | "
236
+ f"{time.time() - t0:.0f}s", flush=True)
237
+ if ep % 4 == 0 or ep == args.epochs:
238
+ print(" sample:", generate(p, dt, "?whale;", seed=ep)[:110], flush=True)
239
+ save(args.save, p, dt, freq, kwc)
240
+ print(f"saved to {args.save}", flush=True)
241
+
242
+
243
+ def evaluate(p, ids, dt, B=50):
244
+ H, N = p["Om"].shape[0], p["Om"].shape[1]
245
+ L = len(ids) // B
246
+ st = ids[: B * L].reshape(B, L)
247
+ x = init_state(B, H, N)
248
+ loss, caches, _ = forward(p, x, st[:, :-1][:, :400], st[:, 1:][:, :400], dt)
249
+ hits = tot = 0
250
+ for cache in caches[50:]: # skip washout
251
+ probs, ys = cache[-3], cache[-1]
252
+ hits += int((probs.argmax(1) == ys).sum())
253
+ tot += len(ys)
254
+ return loss, hits / tot
255
+
256
+
257
+ # ---------------------------------------------------------------- generation
258
+ def step_one(p, x, cid, dt, A):
259
+ x0 = x
260
+ y = couple(p["J"], x0) + p["U"][cid]
261
+ d = (x0 * y).sum(-1, keepdims=True)
262
+ R = np.einsum("inm,bim->bin", A, x0)
263
+ xt = x0 + dt * (R + y - d * x0)
264
+ x = xt / np.sqrt((xt * xt).sum(-1, keepdims=True))
265
+ f, _, _ = readout_feats(p, x)
266
+ return x, f @ p["Wo"] + p["bo"]
267
+
268
+
269
+ def generate(p, dt, prime, max_len=300, temperature=0.8, top_k=8, seed=None):
270
+ rng = np.random.default_rng(seed)
271
+ H, N = p["Om"].shape[0], p["Om"].shape[1]
272
+ A = antisym(p["Om"])
273
+ x = init_state(1, H, N)
274
+ logits = None
275
+ for ch in prime:
276
+ x, logits = step_one(p, x, np.array([C2I.get(ch, 1)]), dt, A)
277
+ out = []
278
+ for _ in range(max_len):
279
+ lg = logits[0] / temperature
280
+ lg -= lg.max()
281
+ pr = np.exp(lg)
282
+ if top_k and top_k < V:
283
+ cut = np.partition(pr, -top_k)[-top_k]
284
+ pr = np.where(pr >= cut, pr, 0.0)
285
+ pr /= pr.sum()
286
+ nxt = int(rng.choice(V, p=pr))
287
+ if CHARS[nxt] == "\n":
288
+ break
289
+ out.append(CHARS[nxt])
290
+ x, logits = step_one(p, x, np.array([nxt]), dt, A)
291
+ return "".join(out)
292
+
293
+
294
+ # ---------------------------------------------------------------- persistence
295
+ def save(path, p, dt, freq, kwc):
296
+ words = np.array(list(freq.keys()))
297
+ counts = np.array(list(freq.values()))
298
+ kwords = np.array(list(kwc.keys()))
299
+ kcounts = np.array(list(kwc.values()))
300
+ np.savez_compressed(path, dt=dt, words=words, counts=counts,
301
+ kwords=kwords, kcounts=kcounts,
302
+ **{f"p_{k}": x for k, x in p.items()})
303
+
304
+
305
+ def load(path):
306
+ zf = np.load(path)
307
+ p = {k[2:]: zf[k] for k in zf.files if k.startswith("p_")}
308
+ freq = dict(zip(zf["words"].tolist(), zf["counts"].tolist()))
309
+ kwc = dict(zip(zf["kwords"].tolist(), zf["kcounts"].tolist()))
310
+ return p, float(zf["dt"]), freq, kwc
311
+
312
+
313
+ # ---------------------------------------------------------------- main
314
+ def main():
315
+ ap = argparse.ArgumentParser(description="KOLM v4 — Kuramoto oscillators")
316
+ ap.add_argument("--text", default="mobydick.txt")
317
+ ap.add_argument("--hidden", type=int, default=256)
318
+ ap.add_argument("--ndim", type=int, default=4, help="oscillator dimension N")
319
+ ap.add_argument("--seq", type=int, default=128)
320
+ ap.add_argument("--batch", type=int, default=64)
321
+ ap.add_argument("--dt", type=float, default=0.25)
322
+ ap.add_argument("--om-init", type=float, default=0.1,
323
+ help="init scale of Om; with dt sets the gradient horizon")
324
+ ap.add_argument("--groups", type=int, default=32,
325
+ help="coherence readout groups G (0 disables)")
326
+ ap.add_argument("--lr", type=float, default=2e-3)
327
+ ap.add_argument("--epochs", type=int, default=18)
328
+ ap.add_argument("--train", action="store_true")
329
+ ap.add_argument("--save", default="kolm256.npz")
330
+ ap.add_argument("--load", default=None)
331
+ ap.add_argument("--ask", default=None, help="one-shot question")
332
+ ap.add_argument("--gradcheck", action="store_true")
333
+ args = ap.parse_args()
334
+
335
+ if args.gradcheck:
336
+ sys.exit(0 if gradcheck() else 1)
337
+ if args.train:
338
+ train(args)
339
+ if not args.ask:
340
+ return
341
+ p, dt, freq, kwc = load(args.load or args.save)
342
+
343
+ def answer(msg):
344
+ kw = pick_keyword(msg, freq, kwc)
345
+ return kw, generate(p, dt, f"?{kw};")
346
+
347
+ if args.ask:
348
+ kw, resp = answer(args.ask)
349
+ print(f"you : {args.ask}\nkolm ({kw}) : {resp}")
350
+ return
351
+ print("KOLM chat — it answers about the topic word of your message (ctrl-d quits)")
352
+ while True:
353
+ try:
354
+ msg = input("\nyou > ")
355
+ except EOFError:
356
+ break
357
+ kw, resp = answer(msg)
358
+ print(f"kolm ({kw}) > {resp}")
359
+
360
+
361
+ if __name__ == "__main__":
362
+ main()