Delon Swartz commited on
add kolm.py
Browse files
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()
|