Skip to content

Instantly share code, notes, and snippets.

@lukaszsamson
Created September 23, 2026 18:11
Show Gist options
  • Select an option

  • Save lukaszsamson/1da96fabce66940cdf49c57544c67b22 to your computer and use it in GitHub Desktop.

Select an option

Save lukaszsamson/1da96fabce66940cdf49c57544c67b22 to your computer and use it in GitHub Desktop.
Port of Karpathy's microgpt to bend2
# The most atomic way to train and inference a GPT in pure, dependency-free Bend.
# This file is the complete algorithm.
# Everything else is just efficiency.
#
# Port of microgpt.py by @karpathy.
#
# Bend is pure and affine, so there is no mutable object graph. A Value is a
# pair V{id, data}; every op appends one node (its children ids and local
# grads) to a tape threaded by a state monad G. Since ids grow with creation
# order, the tape read newest-first is already the reversed topological order
# that `backward` needs, and gradients accumulate in a flat Array<F32>.
import Base
# Model hyperparameters
def n_embd() -> U32:
16
def n_head() -> U32:
4
def n_layer() -> U32:
1
def block_size() -> U32:
16
def head_dim() -> U32:
4
def num_steps() -> U32:
1000
# Randomness: a counter-based generator, so no RNG state is threaded around.
# rand(stream, i) is the i-th draw of an independent stream, seeded with 42.
def hash(+x: U32) -> U32:
+a = (x .^. (x >> 16n) : U32)
+b = (a * 2146121005 : U32)
+c = (b .^. (b >> 15n) : U32)
+d = (c * 2221713035 : U32)
(d .^. (d >> 16n) : U32)
def rand(+stream: U32, +i: U32) -> U32:
hash((hash((stream * 2654435769 + 42 : U32)) + hash(i) : U32))
def uniform(stream: U32, i: U32) -> F32: # in [0, 1)
(U32.to_f32((rand(stream, i) >> 8n : U32)) / 16777216.0 : F32)
def gauss(+i: U32, std: F32) -> F32: # Box-Muller
u1 = (1.0 - uniform(1, (i * 2 : U32)) : F32)
u2 = uniform(1, (i * 2 + 1 : U32))
r = F32.sqrt((F32.neg(2.0) * F32.log(u1) : F32))
(r * F32.cos((2.0 * F32.pi() * u2 : F32)) * std : F32)
# Dataset: lines of input.txt, stripped, empty ones dropped, then shuffled.
def clean.put(empty: Bool, s: String, rest: List<&2, String>) -> List<&2, String>:
match empty:
case True{}:
rest
case False{}:
s <> rest
def clean(xs: List<&2, String>) -> List<&2, String>:
match xs:
case Nil{}:
Nil{}
case h <> t:
+s = String.trim(h)
clean.put(String.is_empty(s), s, clean(t))
type KD is Data:
KD{key: U32, doc: String}
def KD.le(a: KD, b: KD) -> Bool:
match a b:
case KD{x, _} KD{y, _}:
U32.is_le(x, y)
def tag(xs: List<&2, String>, +i: U32) -> List<&2, KD>:
match xs:
case Nil{}:
Nil{}
case h <> t:
KD{rand(2, i), h} <> tag(t, (i + 1 : U32))
def untag(xs: List<&2, KD>) -> List<&2, String>:
match xs:
case Nil{}:
Nil{}
case KD{_, d} <> t:
d <> untag(t)
def shuffle(xs: List<&2, String>) -> List<&2, String>:
untag(List.sort(~KD, ~KD.le, tag(xs, 0)))
# Tokenizer: unique characters become token ids 0..n-1, BOS is n.
def insert.put(o: Cmp, +c: U32, h: U32, t: List<&2, U32>, rest: List<&2, U32>) ->
List<&2, U32>:
match o:
case LT{}:
c <> h <> t
case EQ{}:
h <> t
case GT{}:
h <> rest
def insert(+xs: List<&2, U32>, +c: U32) -> List<&2, U32>:
match xs:
case Nil{}:
[c]
case h <> t:
insert.put(U32.cmp(c, h), c, h, t, insert(t, c))
def chars(cs: List<&2, Char>, acc: List<&2, U32>) -> List<&2, U32>:
match cs:
case Nil{}:
acc
case c <> t:
chars(t, insert(acc, Char.to_u32(c)))
def uniq(docs: List<&2, String>, acc: List<&2, U32>) -> List<&2, U32>:
match docs:
case Nil{}:
acc
case d <> t:
uniq(t, chars(String.to_list(d), acc))
def index.put(hit: Bool, +i: U32, rest: U32) -> U32:
match hit:
case True{}:
i
case False{}:
rest
def index(xs: List<&2, U32>, +c: U32, +i: U32) -> U32:
match xs:
case Nil{}:
i
case h <> t:
index.put(U32.is_eq(h, c), i, index(t, c, (i + 1 : U32)))
def encode(cs: List<&2, Char>, +uchars: List<&2, U32>) -> List<&2, U32>:
match cs:
case Nil{}:
Nil{}
case c <> t:
index(uchars, Char.to_u32(c), 0) <> encode(t, uchars)
def encode_all(docs: List<&2, String>, +uchars: List<&2, U32>) ->
List<&2, List<&2, U32>>:
match docs:
case Nil{}:
Nil{}
case d <> t:
encode(String.to_list(d), uchars) <> encode_all(t, uchars)
# Autograd: a Value, the nodes of the tape, and the tape monad G.
type V is Data:
V{id: U32, data: F32}
type Node is Data:
N0{}
N1{a: U32, la: F32}
N2{a: U32, la: F32, b: U32, lb: F32}
type Tape is Data:
Tape{n: U32, nodes: List<&2, Node>}
law G:
for -A: Data
Type
def G(A):
Tape -> Tape & A
def G.pure(-A: Data, x: A) -> G(A):
t => (t, x)
def G.bind.k(-A: Data, -B: Data, f: A -> G(B), r: Tape & A) -> Tape & B:
(t, x) = r
f(x, t)
def G.bind(-A: Data, -B: Data, m: G(A), f: A -> G(B)) -> G(B):
t => G.bind.k(A, B, f, m(t))
def G.push(nd: Node, d: F32, t: Tape) -> Tape & V:
match t:
case Tape{+n, nodes}:
(Tape{(n + 1 : U32), nd <> nodes}, V{n, d})
def G.node(nd: Node, d: F32) -> G(V):
t => G.push(nd, d, t)
def konst(d: F32) -> G(V):
G.node(N0{}, d)
def add(x: V, y: V) -> G(V):
match x y:
case V{i, a} V{j, b}:
G.node(N2{i, 1.0, j, 1.0}, (a + b : F32))
def mul(x: V, y: V) -> G(V):
match x y:
case V{i, +a} V{j, +b}:
G.node(N2{i, b, j, a}, (a * b : F32))
def scale(x: V, +c: F32) -> G(V): # x * c, for a constant c
match x:
case V{i, a}:
G.node(N1{i, c}, (a * c : F32))
def shift(x: V, c: F32) -> G(V): # x + c, for a constant c
match x:
case V{i, a}:
G.node(N1{i, 1.0}, (a + c : F32))
def vpow(x: V, +e: F32) -> G(V):
match x:
case V{i, +a}:
G.node(N1{i, (e * F32.pow(a, (e - 1.0 : F32)) : F32)}, F32.pow(a, e))
def vlog(x: V) -> G(V):
match x:
case V{i, +a}:
G.node(N1{i, (1.0 / a : F32)}, F32.log(a))
def vexp(x: V) -> G(V):
match x:
case V{i, a}:
+e = F32.exp(a)
G.node(N1{i, e}, e)
def relu.grad(b: Bool) -> F32:
match b:
case True{}:
1.0
case False{}:
0.0
def relu(x: V) -> G(V):
match x:
case V{i, +a}:
G.node(N1{i, relu.grad(F32.is_gt(a, 0.0))}, F32.max(0.0, a))
def V.data(x: V) -> F32:
match x:
case V{_, d}:
d
# Backward: walk the tape newest-first; each visit reads the node's grad
# (zeroing its slot for the next step) and pushes it into its children.
def gadd.fin(r: Array<F32> & F32) -> Array<F32>:
(g, old) = r
g
def gadd(g: Array<F32>, +i: U32, v: F32) -> Array<F32>:
gadd.fin(Array.atomic.fadd(g, i, v))
def chain(nd: Node, r: Array<F32> & F32) -> Array<F32>:
match nd:
case N0{}:
(g, gv) = r
g
case N1{a, la}:
(g, gv) = r
gadd(g, a, (la * gv : F32))
case N2{a, la, b, lb}:
(g, +gv) = r
gadd(gadd(g, a, (la * gv : F32)), b, (lb * gv : F32))
def backward(nodes: List<&2, Node>, g: Array<F32>, +i: U32) -> Array<F32>:
match nodes:
case Nil{}:
g
case nd <> rest:
backward(rest, chain(nd, Array.swap(F32, g, i, 0.0)), (i - 1 : U32))
# Vector helpers, all recording onto the tape.
def vsum.go(xs: List<&2, V>, acc: V) -> G(V):
match xs:
case Nil{}:
G.pure(V, acc)
case h <> t:
do G<V>:
s : V <- add(acc, h)
vsum.go(t, s)
def vsum(xs: List<&2, V>) -> G(V):
match xs:
case Nil{}:
konst(0.0)
case h <> t:
vsum.go(t, h)
def vadd(xs: List<&2, V>, ys: List<&2, V>) -> G(List<&2, V>):
match xs ys:
case Nil{} _:
G.pure(List<&2, V>, Nil{})
case h <> t Nil{}:
G.pure(List<&2, V>, Nil{})
case x <> xt y <> yt:
do G<List<&2, V>>:
z : V <- add(x, y)
zs : List<&2, V> <- vadd(xt, yt)
return z <> zs
def vmul(xs: List<&2, V>, ys: List<&2, V>) -> G(List<&2, V>):
match xs ys:
case Nil{} _:
G.pure(List<&2, V>, Nil{})
case h <> t Nil{}:
G.pure(List<&2, V>, Nil{})
case x <> xt y <> yt:
do G<List<&2, V>>:
z : V <- mul(x, y)
zs : List<&2, V> <- vmul(xt, yt)
return z <> zs
def vmulv(xs: List<&2, V>, +s: V) -> G(List<&2, V>):
match xs:
case Nil{}:
G.pure(List<&2, V>, Nil{})
case x <> t:
do G<List<&2, V>>:
z : V <- mul(x, s)
zs : List<&2, V> <- vmulv(t, s)
return z <> zs
def vscale(xs: List<&2, V>, +c: F32) -> G(List<&2, V>):
match xs:
case Nil{}:
G.pure(List<&2, V>, Nil{})
case x <> t:
do G<List<&2, V>>:
z : V <- scale(x, c)
zs : List<&2, V> <- vscale(t, c)
return z <> zs
def vrelu(xs: List<&2, V>) -> G(List<&2, V>):
match xs:
case Nil{}:
G.pure(List<&2, V>, Nil{})
case x <> t:
do G<List<&2, V>>:
z : V <- relu(x)
zs : List<&2, V> <- vrelu(t)
return z <> zs
def vexp_shift(xs: List<&2, V>, +c: F32) -> G(List<&2, V>): # exp(x + c)
match xs:
case Nil{}:
G.pure(List<&2, V>, Nil{})
case x <> t:
do G<List<&2, V>>:
y : V <- shift(x, c)
z : V <- vexp(y)
zs : List<&2, V> <- vexp_shift(t, c)
return z <> zs
def List.head.v(xs: List<&2, V>) -> V:
match xs:
case Nil{}:
V{0, 0.0}
case h <> t:
h
def maxd(xs: List<&2, V>, acc: F32) -> F32:
match xs:
case Nil{}:
acc
case x <> t:
maxd(t, F32.max(acc, V.data(x)))
def dot(w: List<&2, V>, x: List<&2, V>) -> G(V):
do G<V>:
ps : List<&2, V> <- vmul(w, x)
vsum(ps)
# Model architecture: GPT-2 with rmsnorm, no biases, ReLU.
def linear(w: List<&2, List<&2, V>>, +x: List<&2, V>) -> G(List<&2, V>):
match w:
case Nil{}:
G.pure(List<&2, V>, Nil{})
case row <> rest:
do G<List<&2, V>>:
y : V <- dot(row, x)
ys : List<&2, V> <- linear(rest, x)
return y <> ys
def softmax.norm(e: List<&2, V>) -> G(List<&2, V>):
+ex = e
do G<List<&2, V>>:
total : V <- vsum(ex)
inv : V <- vpow(total, F32.neg(1.0))
vmulv(ex, inv)
def softmax(+logits: List<&2, V>) -> G(List<&2, V>):
G.bind(List<&2, V>, List<&2, V>,
vexp_shift(logits, F32.neg(maxd(logits, V.data(List.head.v(logits))))),
softmax.norm)
def rmsnorm.scale(+x: List<&2, V>, ms: V) -> G(List<&2, V>):
do G<List<&2, V>>:
m : V <- scale(ms, (1.0 / U32.to_f32(n_embd()) : F32))
s : V <- shift(m, 0.00001)
k : V <- vpow(s, F32.neg(0.5))
vmulv(x, k)
def rmsnorm(+x: List<&2, V>) -> G(List<&2, V>):
do G<List<&2, V>>:
ms : V <- dot(x, x)
rmsnorm.scale(x, ms)
def nth(xs: List<&2, List<&2, V>>, n: Nat) -> List<&2, V>:
match xs n:
case Nil{} _:
Nil{}
case h <> t 0n:
h
case h <> t 1n+p:
nth(t, p)
def nth_v(xs: List<&2, V>, n: Nat) -> V:
match xs n:
case Nil{} _:
V{0, 0.0}
case h <> t 0n:
h
case h <> t 1n+p:
nth_v(t, p)
def slice(xs: List<&2, V>, +h: U32) -> List<&2, V>:
List.take(&2, V, List.drop(&2, V, xs, U32.to_nat((h * head_dim() : U32))),
U32.to_nat(head_dim()))
def slices(xss: List<&2, List<&2, V>>, +h: U32) -> List<&2, List<&2, V>>:
match xss:
case Nil{}:
Nil{}
case xs <> t:
slice(xs, h) <> slices(t, h)
def heads_of(xss: List<&2, List<&2, V>>) -> List<&2, V>:
match xss:
case Nil{}:
Nil{}
case xs <> t:
List.head.v(xs) <> heads_of(t)
def tails_of(xss: List<&2, List<&2, V>>) -> List<&2, List<&2, V>>:
match xss:
case Nil{}:
Nil{}
case xs <> t:
List.drop(&2, V, xs, 1n) <> tails_of(t)
def transpose(n: Nat, +m: List<&2, List<&2, V>>) -> List<&2, List<&2, V>>:
match n:
case 0n:
Nil{}
case 1n+p:
heads_of(m) <> transpose(p, tails_of(m))
def head.out(+v_h: List<&2, List<&2, V>>, weights: List<&2, V>) ->
G(List<&2, V>):
linear(transpose(U32.to_nat(head_dim()), v_h), weights)
def head(+h: U32, +q: List<&2, V>, +keys: List<&2, List<&2, V>>,
+vals: List<&2, List<&2, V>>) -> G(List<&2, V>):
do G<List<&2, V>>:
logits : List<&2, V> <- linear(slices(keys, h), slice(q, h))
scaled : List<&2, V> <-
vscale(logits, (1.0 / F32.sqrt(U32.to_f32(head_dim())) : F32))
weights : List<&2, V> <- softmax(scaled)
head.out(slices(vals, h), weights)
def heads(n: Nat, +h: U32, +q: List<&2, V>, +keys: List<&2, List<&2, V>>,
+vals: List<&2, List<&2, V>>) -> G(List<&2, V>):
match n:
case 0n:
G.pure(List<&2, V>, Nil{})
case 1n+p:
do G<List<&2, V>>:
out : List<&2, V> <- head(h, q, keys, vals)
rest : List<&2, V> <- heads(p, (h + 1 : U32), q, keys, vals)
return List.append(&2, V, out, rest)
type Layer is Data:
Layer{
wq: List<&2, List<&2, V>>, wk: List<&2, List<&2, V>>,
wv: List<&2, List<&2, V>>, wo: List<&2, List<&2, V>>,
fc1: List<&2, List<&2, V>>, fc2: List<&2, List<&2, V>>}
type Model is Data:
Model{
wte: List<&2, List<&2, V>>, wpe: List<&2, List<&2, V>>,
lm_head: List<&2, List<&2, V>>, layers: List<&2, Layer>}
# One layer's KV cache: a key and a value vector per position so far.
type KV is Data:
KV{keys: List<&2, List<&2, V>>, vals: List<&2, List<&2, V>>}
type LX is Data:
LX{x: List<&2, V>, kvs: List<&2, KV>}
def LX.x(r: LX) -> List<&2, V>:
match r:
case LX{x, _}:
x
def LX.kvs(r: LX) -> List<&2, KV>:
match r:
case LX{_, kvs}:
kvs
def LX.cons(r: LX, rest: LX) -> LX:
match r rest:
case LX{_, kv} LX{x, kvs}:
LX{x, List.append(&2, KV, kv, kvs)}
def mlp(+x: List<&2, V>, fc1: List<&2, List<&2, V>>,
fc2: List<&2, List<&2, V>>) -> G(List<&2, V>):
do G<List<&2, V>>:
n : List<&2, V> <- rmsnorm(x)
h : List<&2, V> <- linear(fc1, n)
r : List<&2, V> <- vrelu(h)
y : List<&2, V> <- linear(fc2, r)
vadd(y, x)
def block(l: Layer, kv: KV, +x: List<&2, V>) -> G(LX):
match l kv:
case Layer{wq, wk, wv, wo, fc1, fc2} KV{ks, vs}:
do G<LX>:
# 1) Multi-head attention block
+xn : List<&2, V> <- rmsnorm(x)
q : List<&2, V> <- linear(wq, xn)
k : List<&2, V> <- linear(wk, xn)
v : List<&2, V> <- linear(wv, xn)
+keys : List<&2, List<&2, V>> = List.append(&2, List<&2, V>, ks, [k])
+vals : List<&2, List<&2, V>> = List.append(&2, List<&2, V>, vs, [v])
a : List<&2, V> <- heads(U32.to_nat(n_head()), 0, q, keys, vals)
o : List<&2, V> <- linear(wo, a)
x2 : List<&2, V> <- vadd(o, x)
# 2) MLP block
x3 : List<&2, V> <- mlp(x2, fc1, fc2)
return LX{x3, [KV{keys, vals}]}
def blocks(ls: List<&2, Layer>, kvs: List<&2, KV>, x: List<&2, V>) -> G(LX):
match ls kvs:
case Nil{} _:
G.pure(LX, LX{x, Nil{}})
case l <> lt Nil{}:
G.pure(LX, LX{x, Nil{}})
case l <> lt kv <> kt:
G.bind(LX, LX, block(l, kv, x), +r =>
G.bind(LX, LX, blocks(lt, kt, LX.x(r)), rest =>
G.pure(LX, LX.cons(r, rest))))
def gpt.head(lm_head: List<&2, List<&2, V>>, r: LX) -> G(LX):
match r:
case LX{x, kvs}:
do G<LX>:
logits : List<&2, V> <- linear(lm_head, x)
return LX{logits, kvs}
# gpt answers the logits in LX.x, and the grown caches in LX.kvs
def gpt(+tok: U32, +pos: U32, m: Model, kvs: List<&2, KV>) -> G(LX):
match m:
case Model{wte, wpe, lm_head, layers}:
do G<LX>:
x : List<&2, V> <-
vadd(nth(wte, U32.to_nat(tok)), nth(wpe, U32.to_nat(pos)))
+xn : List<&2, V> <- rmsnorm(x)
r : LX <- blocks(layers, kvs, xn)
gpt.head(lm_head, r)
def empty_kvs(n: Nat) -> List<&2, KV>:
match n:
case 0n:
Nil{}
case 1n+p:
KV{Nil{}, Nil{}} <> empty_kvs(p)
# Parameters, and building the model's matrices out of them.
type PS is Data: # a parameter with its Adam buffers
PS{p: F32, m: F32, v: F32}
def init_params(n: Nat, +i: U32) -> List<&2, PS>:
match n:
case 0n:
Nil{}
case 1n+p:
PS{gauss(i, 0.08), 0.0, 0.0} <> init_params(p, (i + 1 : U32))
def values(ps: List<&2, PS>, +i: U32) -> List<&2, V>:
match ps:
case Nil{}:
Nil{}
case PS{p, _, _} <> t:
V{i, p} <> values(t, (i + 1 : U32))
def rows(n: Nat, xs: List<&2, V>, +c: Nat) -> List<&2, List<&2, V>>:
match n:
case 0n:
Nil{}
case 1n+p:
+ys = xs
List.take(&2, V, ys, c) <> rows(p, List.drop(&2, V, ys, c), c)
def matrix(+vs: List<&2, V>, off: U32, +nout: U32, +nin: U32) ->
List<&2, List<&2, V>>:
rows(U32.to_nat(nout),
List.drop(&2, V, vs, U32.to_nat(off)), U32.to_nat(nin))
def layer_size() -> U32:
(12 * n_embd() * n_embd() : U32)
def layers(n: Nat, +vs: List<&2, V>, +off: U32) -> List<&2, Layer>:
match n:
case 0n:
Nil{}
case 1n+p:
+e = n_embd()
+s = (e * e : U32)
Layer{
matrix(vs, off, e, e),
matrix(vs, (off + s : U32), e, e),
matrix(vs, (off + 2 * s : U32), e, e),
matrix(vs, (off + 3 * s : U32), e, e),
matrix(vs, (off + 4 * s : U32), (4 * e : U32), e),
matrix(vs, (off + 8 * s : U32), e, (4 * e : U32))}
<> layers(p, vs, (off + layer_size() : U32))
def num_params(+vocab: U32) -> U32:
(2 * vocab * n_embd() + block_size() * n_embd() + n_layer() * layer_size() : U32)
def model(+vocab: U32, ps: List<&2, PS>) -> Model:
+vs = values(ps, 0)
+e = n_embd()
+b = block_size()
Model{
matrix(vs, 0, vocab, e),
matrix(vs, (vocab * e : U32), b, e),
matrix(vs, (vocab * e + b * e : U32), vocab, e),
layers(U32.to_nat(n_layer()), vs, (2 * vocab * e + b * e : U32))}
# Training: forward the document, build up the tape all the way to the loss.
type TT is Data:
TT{tok: U32, tgt: U32}
def pairs(n: Nat, toks: List<&2, U32>) -> List<&2, TT>:
match n toks:
case 0n _:
Nil{}
case 1n+p Nil{}:
Nil{}
case 1n+p a <> Nil{}:
Nil{}
case 1n+p a <> +b <> t:
TT{a, b} <> pairs(p, b <> t)
def forward(ps: List<&2, TT>, +pos: U32, +m: Model, kvs: List<&2, KV>,
losses: List<&2, V>) -> G(List<&2, V>):
match ps:
case Nil{}:
G.pure(List<&2, V>, losses)
case TT{tok, tgt} <> rest:
G.bind(LX, List<&2, V>, gpt(tok, pos, m, kvs), +r =>
G.bind(List<&2, V>, List<&2, V>, softmax(LX.x(r)), probs =>
G.bind(V, List<&2, V>, vlog(nth_v(probs, U32.to_nat(tgt))), lp =>
forward(rest, (pos + 1 : U32), m, LX.kvs(r), lp <> losses))))
def loss(ps: List<&2, TT>, +n: U32, +m: Model) -> G(V):
do G<V>:
logps : List<&2, V> <- forward(ps, 0, m, empty_kvs(U32.to_nat(n_layer())), Nil{})
total : V <- vsum(logps)
scale(total, (F32.neg(1.0) / U32.to_f32(n) : F32)) # mean of -log p
# Adam, the blessed optimizer
type Hyp is Data:
Hyp{lr: F32, c1: F32, c2: F32}
def adam1(p: PS, +g: F32, +h: Hyp) -> PS:
match p h:
case PS{w, m, v} Hyp{lr, c1, c2}:
+m2 = (0.85 * m + 0.15 * g : F32)
+v2 = (0.99 * v + 0.01 * g * g : F32)
+mh = (m2 / c1 : F32)
+vh = (v2 / c2 : F32)
PS{(w - lr * mh / (F32.sqrt(vh) + 0.00000001) : F32), m2, v2}
type St is Type:
St{ps: List<&2, PS>, grads: Array<F32>}
def adam(ps: List<&2, PS>, r: Array<F32> & F32, +i: U32, +h: Hyp,
acc: List<&2, PS>) -> St:
match ps:
case Nil{}:
(g, _) = r
St{List.reverse(&2, PS, acc), g}
case p <> t:
(g, gi) = r # the param's gradient, its slot zeroed for the next step
+j = (i + 1 : U32)
adam(t, Array.swap(F32, g, j, 0.0), j, h, adam1(p, gi, h) <> acc)
def hyp(+step: U32) -> Hyp:
+t = U32.to_f32((step + 1 : U32))
lr = (0.01 * (1.0 - U32.to_f32(step) / U32.to_f32(num_steps())) : F32)
Hyp{lr, (1.0 - F32.pow(0.85, t) : F32), (1.0 - F32.pow(0.99, t) : F32)}
type Cfg is Data:
Cfg{uchars: List<&2, U32>, vocab: U32, ndocs: U32}
def Cfg.vocab(c: Cfg) -> U32:
match c:
case Cfg{_, v, _}:
v
def doc_at(docs: List<&2, List<&2, U32>>, n: Nat) -> List<&2, U32>:
match docs n:
case Nil{} _:
Nil{}
case h <> t 0n:
h
case h <> t 1n+p:
doc_at(t, p)
def train.update(+step: U32, ps: List<&2, PS>, g: Array<F32>, r: Tape & V) ->
St & F32:
(t, l) = r
match t:
case Tape{+top, nodes}:
+last = (top - 1 : U32)
g2 = backward(nodes, Array.set(F32, g, last, 1.0), last)
(adam(ps, Array.swap(F32, g2, 0, 0.0), 0, hyp(step), Nil{}), V.data(l))
def train_step(+step: U32, +cfg: Cfg, +docs: List<&2, List<&2, U32>>, st: St) ->
St & F32:
match cfg st:
case Cfg{_, +vocab, +ndocs} St{+ps, g}:
# Take a single document, tokenize it, surround it with BOS on both sides
doc = doc_at(docs, U32.to_nat((step % ndocs : U32)))
+bos = (vocab - 1 : U32)
+toks = {bos <> List.append(&2, U32, doc, [bos]) : List<&2, U32>}
+n = U32.min(block_size(), (U32.from_nat(List.length(&2, U32, toks)) - 1 : U32))
+m = model(vocab, ps)
tape = {Tape{num_params(vocab), Nil{}} : Tape}
train.update(step, ps, g, loss(pairs(U32.to_nat(n), toks), n, m)(tape))
# Output formatting
def pad(s: String, n: Nat) -> String: # left-pad s with spaces to n chars
+t = s
String.repeat(" ", Nat.sub(n, String.length(t))) ++ t
def fixed4(+x: F32) -> String: # x with 4 decimals, for x >= 0
+k = F32.to_u32(F32.round((x * 10000.0 : F32)))
+frac = U32.show((k % 10000 : U32))
U32.show((k / 10000 : U32)) ++ "." ++
String.repeat("0", Nat.sub(4n, String.length(frac))) ++ frac
def show_step(+step: U32, l: F32) -> String:
"step " ++ pad(U32.show(step), 4n) ++ " / " ++ pad(U32.show(num_steps()), 4n)
++ " | loss " ++ fixed4(l)
def train(n: Nat, r: St & F32, +step: U32, +cfg: Cfg,
+docs: List<&2, List<&2, U32>>) -> IO(St):
match n:
case 0n:
(st, l) = r
do IO<St>:
IO.print(show_step(step, l))
return st
case 1n+p:
(st, l) = r
do IO<St>:
IO.print(show_step(step, l))
train(p, train_step(step, cfg, docs, st), (step + 1 : U32), cfg, docs)
# Inference: may the model babble back to us
def sm_exps(xs: List<&2, V>, +mx: F32, +temp: F32) -> List<&2, F32>:
match xs:
case Nil{}:
Nil{}
case x <> t:
F32.exp(((V.data(x) - mx) / temp : F32)) <> sm_exps(t, mx, temp)
def sumf(xs: List<&2, F32>, acc: F32) -> F32:
match xs:
case Nil{}:
acc
case x <> t:
sumf(t, (acc + x : F32))
def choose.put(hit: Bool, +i: U32, rest: U32) -> U32:
match hit:
case True{}:
i
case False{}:
rest
def choose(ws: List<&2, F32>, +r: F32, +i: U32) -> U32:
match ws:
case Nil{}:
(i - 1 : U32)
case +w <> t:
choose.put(F32.is_lt(r, w), i, choose(t, (r - w : F32), (i + 1 : U32)))
def pick(+logits: List<&2, V>, u: F32) -> U32: # random.choices on softmax(logits / T)
+ws = sm_exps(logits, maxd(logits, V.data(List.head.v(logits))), 0.5)
choose(ws, (u * sumf(ws, 0.0) : F32), 0)
def tail(xs: List<&2, U32>) -> List<&2, U32>:
match xs:
case Nil{}:
Nil{}
case h <> t:
t
def sample(n: Nat, done: Bool, +tok: U32, +pos: U32, +m: Model,
kvs: List<&2, KV>, acc: List<&2, U32>, +idx: U32, +bos: U32) ->
G(List<&2, U32>):
match n done:
case 0n False{}:
G.pure(List<&2, U32>, acc)
case 0n True{}:
G.pure(List<&2, U32>, tail(acc))
case 1n+p True{}:
G.pure(List<&2, U32>, tail(acc))
case 1n+p False{}:
G.bind(LX, List<&2, U32>, gpt(tok, pos, m, kvs), +r =>
+next = pick(LX.x(r), uniform(3, (idx * block_size() + pos : U32)))
sample(p, U32.is_eq(next, bos), next, (pos + 1 : U32), m, LX.kvs(r),
next <> acc, idx, bos))
def U32.nth(xs: List<&2, U32>, n: Nat) -> U32:
match xs n:
case Nil{} _:
0
case h <> t 0n:
h
case h <> t 1n+p:
U32.nth(t, p)
def decode(toks: List<&2, U32>, +uchars: List<&2, U32>, acc: String) -> String:
match toks:
case Nil{}:
acc
case t <> rest:
decode(rest, uchars,
SCon{Char.from_u32(U32.nth(uchars, U32.to_nat(t))), acc})
def sample.text(+uchars: List<&2, U32>, r: Tape & List<&2, U32>) -> String:
(t, toks) = r
decode(toks, uchars, "") # toks are newest-first, so decode prepends
def samples(n: Nat, +i: U32, +cfg: Cfg, +m: Model) -> IO(Unit):
match n:
case 0n:
IO.pure(Unit, Unit{})
case 1n+p:
match cfg:
case Cfg{+uchars, +vocab, +ndocs}:
tape = {Tape{num_params(vocab), Nil{}} : Tape}
+bos = (vocab - 1 : U32)
text = sample.text(uchars, sample(U32.to_nat(block_size()), False{},
bos, 0, m, empty_kvs(U32.to_nat(n_layer())), Nil{}, i, bos)(tape))
do IO<Unit>:
IO.print("sample " ++ pad(U32.show((i + 1 : U32)), 2n) ++ ": " ++ text)
samples(p, (i + 1 : U32), Cfg{uchars, vocab, ndocs}, m)
# The main program
def run.infer(+cfg: Cfg, st: St) -> IO(Unit):
match st:
case St{ps, g}:
do IO<Unit>:
IO.print("")
IO.print("--- inference (new, hallucinated names) ---")
samples(20n, 0, cfg, model(Cfg.vocab(cfg), ps))
def run(text: String) -> IO(Unit):
+docs = shuffle(clean(String.lines(text)))
+uchars = uniq(docs, Nil{})
+vocab = (U32.from_nat(List.length(&2, U32, uchars)) + 1 : U32)
+ndocs = U32.from_nat(List.length(&2, String, docs))
+cfg = {Cfg{uchars, vocab, ndocs} : Cfg}
+np = num_params(vocab)
+toks = encode_all(docs, uchars)
# grads: one slot per tape node; a 16-token doc makes ~120k nodes
st = {St{init_params(U32.to_nat(np), 0), [0.0 : F32^19n]} : St}
do IO<Unit>:
IO.print("num docs: " ++ U32.show(ndocs))
IO.print("vocab size: " ++ U32.show(vocab))
IO.print("num params: " ++ U32.show(np))
fin : St <- train(U32.to_nat((num_steps() - 1 : U32)),
train_step(0, cfg, toks, st), 1, cfg, toks)
run.infer(cfg, fin)
def main.read(m: File & Result<&1, &1, U32 & String, String>) -> IO(Unit):
(f, r) = m
do IO<Unit>:
text : String <- IO.pass(String, r)
x : Unit <- File.close(f)
run(text)
def main() -> IO(Unit):
do IO<Unit>:
f : File <- IO.try(File, File.open("input.txt", "r"))
m : File & Result<&1, &1, U32 & String, String> <- File.read(f, 16777216)
main.read(m)
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment