~/bend-docscommunity

main.bend checks

raw source on the hub · import bend-ml-bpe-tokenizer@0.1.2.0/main.bend as Main

bend-ml-bpe-tokenizer: byte-level BPE with a proved roundtrip.

import bend-ml-bpe-tokenizer@0.1.2.0/main.bend as BPE

A token is a raw byte (B{n}) or the result of a merge rule (M{id}). The table is a list of rules from the oldest to the newest, like GPT-2's merges.txt: the rule with id k joins the pair (a, b) into the token M{k}. encode applies the rules in order; decode expands each token back into bytes. The laws are at the end of the file, each with its proof.

How to read the proofs: a proof in Bend is a function whose TYPE is the statement. match does case analysis, the recursive call is the induction hypothesis, %e : P rewrites the goal with the equality e, {==} closes when the two sides are already the same term. +x marks a variable that may be used more than once.

2 imports
import Base
import bend-ml-nat-lemmas@0.1.0.0/main.bend as NL

Laws

law roundtrip provedsource · line 514 · raw

@+table:List<&2, Rule> -> @+bs:List<&2, Nat> -> @+w:{True{} == wf(table) : Bool} -> {decode(table, encode(table, lift(bs))) == bs : List<&2, Nat>}

LAW 1 (roundtrip): with a well-formed table (all rule ids distinct), decoding what was encoded returns the original bytes.

law vocab_bound provedsource · line 525 · raw

@+table:List<&2, Rule> -> @+bs:List<&2, Nat> -> @+w:{True{} == wf(table) : Bool} -> {True{} == knl(List.reverse(&2, Rule, table), encode(table, lift(bs))) : Bool}

LAW 2 (vocabulary): every token emitted by encode is a raw byte or was created by one of the table's rules; an unknown id never appears.

law dec_append provedsource · line 536 · raw

@+D:List<&2, Rule> -> @xs:List<&2, Tk> -> @+ys:List<&2, Tk> -> {dec(D, List.append(&2, Tk, xs, ys)) == List.append(&2, Nat, dec(D, xs), dec(D, ys)) : List<&2, Nat>}

LAW 3 (concatenation): decoding two token lists together is decoding each one and joining the bytes.

law train_wf provedsource · line 638 · raw

@n:Nat -> @+next:Nat -> @+xs:List<&2, Tk> -> {True{} == wf(train(n, next, xs)) : Bool}

LAW (train_wf): the table that train returns is always well formed, so the roundtrip holds for it

law roundtrip_trained provedsource · line 649 · raw

@+n:Nat -> @+corpus:List<&2, Nat> -> @+bs:List<&2, Nat> -> {decode(train(n, 0n, lift(corpus)), encode(train(n, 0n, lift(corpus)), lift(bs))) == bs : List<&2, Nat>}

LAW 5 (roundtrip_trained): train a table on any corpus, with as many merges as you like: decoding what it encodes always returns the original bytes. No hypothesis at all.

Types

type Tk source · line 21 · raw

Data

Token: a raw byte (B) or the result of a merge rule, by rule id (M).

type Rule source · line 26 · raw

Data

Merge rule: joins the pair (a, b) into the token M{id}.

type Cnt source · line 177 · raw

Data

===================================================================== Training (needs no proof: it generates tables; a table is valid if wf(table)) ===================================================================== --------------------------------------------------------------- training: count adjacent pairs, merge the most frequent one, repeat (tie-break: the first pair that appeared, as in minbpe) ---------------------------------------------------------------

Definitions

def Tk.eq source · line 29 · raw

@+x:Tk -> @+y:Tk -> Bool

def peek source · line 41 · raw

@+a:Tk -> @+b:Tk -> @xs:List<&2, Tk> -> Bool

does xs start with exactly the pair (a, b)?

def merge.go source · line 55 · raw

@xs:List<&2, Tk> -> @hit:Bool -> @+a:Tk -> @+b:Tk -> @+c:Tk -> List<&2, Tk>

Replaces each occurrence (non-overlapping, left to right) of the pair (a, b) with the token c. hit is always peek(a, b, xs): the caller computes it, because a match can only open parameters.

def merge source · line 66 · raw

@+a:Tk -> @+b:Tk -> @+c:Tk -> @+xs:List<&2, Tk> -> List<&2, Tk>

def encode.go source · line 71 · raw

@todo:List<&2, Rule> -> @done:List<&2, Rule> -> @xs:List<&2, Tk> -> List<&2, Tk>

applies the rules from the oldest to the newest; done keeps the ones already applied (newest first), which is exactly the table that decode needs.

def encode source · line 81 · raw

@+table:List<&2, Rule> -> @xs:List<&2, Tk> -> List<&2, Tk>

The table goes from the oldest rule to the newest (like merges.txt).

def hits source · line 85 · raw

@table:List<&2, Rule> -> @+tk:Tk -> Bool

is the newest rule in the table the one that creates the token tk?

def exp.go source · line 99 · raw

@table:List<&2, Rule> -> @hit:Bool -> @tk:Tk -> List<&2, Nat>

Expansion of a token into bytes. The table has the newest rule first and a rule only refers to older tokens, so expanding a and b looks only at the rest of the table. hit is always hits(table, tk).

def exp source · line 112 · raw

@+table:List<&2, Rule> -> @+tk:Tk -> List<&2, Nat>

def dec source · line 116 · raw

@+table:List<&2, Rule> -> @ids:List<&2, Tk> -> List<&2, Nat>

decodes with the already reversed table (newest first)

def decode source · line 123 · raw

@+table:List<&2, Rule> -> @ids:List<&2, Tk> -> List<&2, Nat>

def lift source · line 127 · raw

@bs:List<&2, Nat> -> List<&2, Tk>

bytes -> initial tokens

def has source · line 137 · raw

@+D:List<&2, Rule> -> @+i:Nat -> Bool

does some rule in the table have the id i?

def kn source · line 145 · raw

@+D:List<&2, Rule> -> @tk:Tk -> Bool

is the token a byte, or was it created by a rule in the table?

def knl source · line 152 · raw

@+D:List<&2, Rule> -> @xs:List<&2, Tk> -> Bool

def wf.go source · line 160 · raw

@todo:List<&2, Rule> -> @+done:List<&2, Rule> -> Bool

well-formed table: no id repeats

def wf source · line 167 · raw

@table:List<&2, Rule> -> Bool

def same_head source · line 181 · raw

@cs:List<&2, Cnt> -> @+a:Tk -> @+b:Tk -> Bool

is the first Cnt of cs the pair (a, b)?

def bump.go source · line 189 · raw

@cs:List<&2, Cnt> -> @hit:Bool -> @+a:Tk -> @+b:Tk -> List<&2, Cnt>

adds 1 to the counter of the pair (a, b), or creates it at the end of the list

def bump source · line 198 · raw

@+cs:List<&2, Cnt> -> @+a:Tk -> @+b:Tk -> List<&2, Cnt>

def count.go source · line 201 · raw

@xs:List<&2, Tk> -> @cs:List<&2, Cnt> -> List<&2, Cnt>

def pick2 source · line 212 · raw

@better:Bool -> @c:Cnt -> @cur:Cnt -> Cnt

def pick source · line 219 · raw

@c:Cnt -> @cur:Cnt -> Cnt

def best.go source · line 224 · raw

@cs:List<&2, Cnt> -> @cur:Cnt -> Cnt

def best source · line 231 · raw

@cs:List<&2, Cnt> -> Maybe<&1, Cnt>

def train.go source · line 239 · raw

@n:Nat -> @+next:Nat -> @+xs:List<&2, Tk> -> @m:Maybe<&1, Cnt> -> List<&2, Rule>

up to n merges, ids from next on; m is the current best pair (best(count(xs)))

def train source · line 250 · raw

@n:Nat -> @+next:Nat -> @+xs:List<&2, Tk> -> List<&2, Rule>

the table from the oldest rule to the newest

def BD source · line 258 · raw

@b:Bool -> @t:Type -> @f:Type -> Type

a type chosen by a Bool: rewriting through it refutes True == False

def nat_eq_sound source · line 265 · raw

@a:Nat -> @b:Nat -> @e:{True{} == Nat.is_eq(a, b) : Bool} -> {a == b : Nat}

def nat_eq_refl source · line 279 · raw

@a:Nat -> {True{} == Nat.is_eq(a, a) : Bool}

def and_true source · line 286 · raw

@p:Bool -> @q:Bool -> @e:{True{} == Bool.and(p, q) : Bool} -> Pair({True{} == p : Bool}, {True{} == q : Bool})

def and_intro source · line 294 · raw

@p:Bool -> @q:Bool -> @ep:{True{} == p : Bool} -> @eq:{True{} == q : Bool} -> {True{} == Bool.and(p, q) : Bool}

def or_left source · line 302 · raw

@p:Bool -> @q:Bool -> @ep:{True{} == p : Bool} -> {True{} == Bool.or(p, q) : Bool}

def or_right source · line 310 · raw

@p:Bool -> @q:Bool -> @eq:{True{} == q : Bool} -> {True{} == Bool.or(p, q) : Bool}

def not_true source · line 318 · raw

@p:Bool -> @e:{True{} == Bool.not(p) : Bool} -> {False{} == p : Bool}

not p is true -> p is false

def tk_eq_sound source · line 326 · raw

@+x:Tk -> @+y:Tk -> @e:{True{} == Tk.eq(x, y) : Bool} -> {x == y : Tk}

def ext_m source · line 343 · raw

@+D:List<&2, Rule> -> @+i:Nat -> @+a:Tk -> @+b:Tk -> @+j:Nat -> @h:Bool -> @eh:{h == Nat.is_eq(i, j) : Bool} -> @nf:{True{} == Bool.not(has(D, i)) : Bool} -> @kt:{True{} == has(D, j) : Bool} -> {exp(Rule{i, a, b} <> D, M{j}) == exp(D, M{j}) : List<&2, Nat>}

Adding the (new) rule i on top of the table does not change the expansion of a token M{j} that the old table already knows (j != i, because i is new).

def and_true_l source · line 357 · raw

@p:Bool -> @q:Bool -> @e:{True{} == Bool.and(p, q) : Bool} -> {True{} == p : Bool}

def and_true_r source · line 365 · raw

@p:Bool -> @q:Bool -> @e:{True{} == Bool.and(p, q) : Bool} -> {True{} == q : Bool}

def ext_tok source · line 374 · raw

@+D:List<&2, Rule> -> @+i:Nat -> @+a:Tk -> @+b:Tk -> @t:Tk -> @kt:{True{} == kn(D, t) : Bool} -> @nf:{True{} == Bool.not(has(D, i)) : Bool} -> {exp(Rule{i, a, b} <> D, t) == exp(D, t) : List<&2, Nat>}

the same for any known token (a byte never depends on the table)

def ext_dec.step source · line 383 · raw

@+D:List<&2, Rule> -> @+i:Nat -> @+a:Tk -> @+b:Tk -> @+x:Tk -> @+t:List<&2, Tk> -> @k1:{True{} == kn(D, x) : Bool} -> @nf:{True{} == Bool.not(has(D, i)) : Bool} -> @ih:{dec(Rule{i, a, b} <> D, t) == dec(D, t) : List<&2, Nat>} -> {dec(Rule{i, a, b} <> D, x <> t) == dec(D, x <> t) : List<&2, Nat>}

def ext_dec source · line 389 · raw

@+D:List<&2, Rule> -> @+i:Nat -> @+a:Tk -> @+b:Tk -> @xs:List<&2, Tk> -> @+kx:{True{} == knl(D, xs) : Bool} -> @+nf:{True{} == Bool.not(has(D, i)) : Bool} -> {dec(Rule{i, a, b} <> D, xs) == dec(D, xs) : List<&2, Nat>}

decoding a list of known tokens does not change when the new rule is added

def exp_new source · line 398 · raw

@+D:List<&2, Rule> -> @+i:Nat -> @+a:Tk -> @+b:Tk -> @h:Bool -> @eh:{h == Nat.is_eq(i, i) : Bool} -> {exp(Rule{i, a, b} <> D, M{i}) == List.append(&2, Nat, exp(D, a), exp(D, b)) : List<&2, Nat>}

The new token M{i} expands to the expansion of a followed by that of b (looking only at the old table D).

def mp_core source · line 410 · raw

@+D:List<&2, Rule> -> @+i:Nat -> @+a:Tk -> @+b:Tk -> @+t:List<&2, Tk> -> @+m:List<&2, Tk> -> @ka:{True{} == kn(D, a) : Bool} -> @kb:{True{} == kn(D, b) : Bool} -> @+nf:{True{} == Bool.not(has(D, i)) : Bool} -> @ihm:{dec(Rule{i, a, b} <> D, m) == dec(Rule{i, a, b} <> D, t) : List<&2, Nat>} -> {dec(Rule{i, a, b} <> D, M{i} <> m) == dec(Rule{i, a, b} <> D, a <> b <> t) : List<&2, Nat>}

Central step: in place of a b t we get M{i} followed by m, where m decodes like t. Both sides decode the same because M{i} expands to exp(a) ++ exp(b).

def kn_sub source · line 418 · raw

@+D:List<&2, Rule> -> @+x:Tk -> @+a:Tk -> @ex:{x == a : Tk} -> @kx:{True{} == kn(D, x) : Bool} -> {True{} == kn(D, a) : Bool}

if x == a and x is known, a is known too

def mp_false source · line 423 · raw

@+D:List<&2, Rule> -> @+i:Nat -> @+a:Tk -> @+b:Tk -> @+x:Tk -> @+t:List<&2, Tk> -> @+mg:List<&2, Tk> -> @ih:{dec(Rule{i, a, b} <> D, mg) == dec(Rule{i, a, b} <> D, t) : List<&2, Nat>} -> {dec(Rule{i, a, b} <> D, x <> mg) == dec(Rule{i, a, b} <> D, x <> t) : List<&2, Nat>}

x stays out of the rewrite: if the rest decodes the same, x <> rest does too

def mp_hit source · line 428 · raw

@+D:List<&2, Rule> -> @+i:Nat -> @+a:Tk -> @+b:Tk -> @+x:Tk -> @+y:Tk -> @+t:List<&2, Tk> -> @+m:List<&2, Tk> -> @+ex:{x == a : Tk} -> @+ey:{y == b : Tk} -> @+kxx:{True{} == kn(D, x) : Bool} -> @+kyy:{True{} == kn(D, y) : Bool} -> @+nf:{True{} == Bool.not(has(D, i)) : Bool} -> @ihm:{dec(Rule{i, a, b} <> D, m) == dec(Rule{i, a, b} <> D, t) : List<&2, Nat>} -> {dec(Rule{i, a, b} <> D, M{i} <> m) == dec(Rule{i, a, b} <> D, x <> y <> t) : List<&2, Nat>}

the pair is x y: the tokens were equal to a and b

def mp source · line 434 · raw

@+D:List<&2, Rule> -> @+i:Nat -> @+a:Tk -> @+b:Tk -> @xs:List<&2, Tk> -> @hit:Bool -> @+eh:{hit == peek(a, b, xs) : Bool} -> @+kx:{True{} == knl(D, xs) : Bool} -> @+nf:{True{} == Bool.not(has(D, i)) : Bool} -> {dec(Rule{i, a, b} <> D, merge.go(xs, hit, a, b, M{i})) == dec(Rule{i, a, b} <> D, xs) : List<&2, Nat>}

merge.go does not change what decode returns (with the new rule on top of the table)

def kn_mono source · line 449 · raw

@+D:List<&2, Rule> -> @+i:Nat -> @+a:Tk -> @+b:Tk -> @t:Tk -> @kt:{True{} == kn(D, t) : Bool} -> {True{} == kn(Rule{i, a, b} <> D, t) : Bool}

adding a rule to the table only increases what is known

def kn_new source · line 457 · raw

@+D:List<&2, Rule> -> @+i:Nat -> @+a:Tk -> @+b:Tk -> {True{} == kn(Rule{i, a, b} <> D, M{i}) : Bool}

the token created by the new rule is known

def kp source · line 460 · raw

@+D:List<&2, Rule> -> @+i:Nat -> @+a:Tk -> @+b:Tk -> @xs:List<&2, Tk> -> @hit:Bool -> @+eh:{hit == peek(a, b, xs) : Bool} -> @+kx:{True{} == knl(D, xs) : Bool} -> {True{} == knl(Rule{i, a, b} <> D, merge.go(xs, hit, a, b, M{i})) : Bool}

def enc_known source · line 475 · raw

@todo:List<&2, Rule> -> @+done:List<&2, Rule> -> @+xs:List<&2, Tk> -> @+wfe:{True{} == wf.go(todo, done) : Bool} -> @+kx:{True{} == knl(done, xs) : Bool} -> {True{} == knl(List.reverse.go(&2, Rule, todo, done), encode.go(todo, done, xs)) : Bool}

every emitted token is known to the complete (reversed) table

def enc_dec source · line 483 · raw

@todo:List<&2, Rule> -> @+done:List<&2, Rule> -> @+xs:List<&2, Tk> -> @+wfe:{True{} == wf.go(todo, done) : Bool} -> @+kx:{True{} == knl(done, xs) : Bool} -> {dec(List.reverse.go(&2, Rule, todo, done), encode.go(todo, done, xs)) == dec(done, xs) : List<&2, Nat>}

applying all the rules does not change what decode returns

def lift_dec source · line 493 · raw

@bs:List<&2, Nat> -> {dec([], lift(bs)) == bs : List<&2, Nat>}

with no rules, the bytes decode to themselves

def lift_known source · line 501 · raw

@bs:List<&2, Nat> -> {True{} == knl([], lift(bs)) : Bool}

def lt_succ source · line 560 · raw

@n:Nat -> {True{} == Nat.is_lt(n, 1n+n) : Bool}

n < 1 + n

def lt_ne source · line 568 · raw

@j:Nat -> @i:Nat -> @e:{True{} == Nat.is_lt(j, i) : Bool} -> {False{} == Nat.is_eq(i, j) : Bool}

j < i -> i != j

def lt_up source · line 582 · raw

@j:Nat -> @i:Nat -> @e:{True{} == Nat.is_lt(j, i) : Bool} -> {True{} == Nat.is_lt(j, 1n+i) : Bool}

j < i -> j < 1 + i

def below source · line 596 · raw

@D:List<&2, Rule> -> @+i:Nat -> Bool

all ids in the rule list are smaller than i

def below_has source · line 603 · raw

@D:List<&2, Rule> -> @+i:Nat -> @+e:{True{} == below(D, i) : Bool} -> {False{} == has(D, i) : Bool}

def below_up source · line 611 · raw

@D:List<&2, Rule> -> @+i:Nat -> @+e:{True{} == below(D, i) : Bool} -> {True{} == below(D, 1n+i) : Bool}

def below_cons source · line 619 · raw

@+a:Tk -> @+b:Tk -> @+D:List<&2, Rule> -> @+i:Nat -> @+e:{True{} == below(D, i) : Bool} -> {True{} == below(Rule{i, a, b} <> D, 1n+i) : Bool}

adding the rule with id next (greater than all those in done) keeps the table "below 1+next"

def below_fresh source · line 623 · raw

@+D:List<&2, Rule> -> @+i:Nat -> @+e:{True{} == below(D, i) : Bool} -> {True{} == Bool.not(has(D, i)) : Bool}

not (has D i) when all the ids in D are smaller than i

def tw source · line 628 · raw

@n:Nat -> @+next:Nat -> @+xs:List<&2, Tk> -> @m:Maybe<&1, Cnt> -> @+done:List<&2, Rule> -> @+bl:{True{} == below(done, next) : Bool} -> {True{} == wf.go(train.go(n, next, xs, m), done) : Bool}

train.go generates rules with ids next, next+1, ...: none repeats, whatever the corpus