Created
September 23, 2026 18:11
-
-
Save lukaszsamson/1da96fabce66940cdf49c57544c67b22 to your computer and use it in GitHub Desktop.
Port of Karpathy's microgpt to bend2
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| # 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