diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-09-11 18:38:13 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-09-11 18:38:13 +0200 |
| commit | 75e581b3bc19a73d0c1f48b32f2e58121c9741c1 (patch) | |
| tree | 897e2f24db3374d58ba36dcee67879cc7cd4956c | |
| parent | 330c6387962b197c5ca7a8051d4864feaecc0f25 (diff) | |
| parent | f998fee7427403698f9262d1e92d1a307c652094 (diff) | |
Merge remote-tracking branch 'origin/wip-entropy' into wip-frequency
| -rw-r--r-- | dict/dict.go | 89 | ||||
| -rw-r--r-- | dict/dict_test.go | 78 | ||||
| -rw-r--r-- | dict/go.mod | 3 | ||||
| -rw-r--r-- | dict/serialize.go | 55 | ||||
| -rw-r--r-- | go.work | 2 | ||||
| -rw-r--r-- | research/entropy/cmd/segment/main.go | 134 | ||||
| -rw-r--r-- | research/entropy/context.go | 53 | ||||
| -rw-r--r-- | research/entropy/entry.go | 9 | ||||
| -rw-r--r-- | research/entropy/experiment.go | 210 | ||||
| -rw-r--r-- | research/entropy/go.mod | 32 | ||||
| -rw-r--r-- | research/entropy/go.sum | 36 | ||||
| -rw-r--r-- | tokenizer/byte/tokenizer.go | 94 | ||||
| -rw-r--r-- | tokenizer/byte/tokenizer_test.go | 49 |
13 files changed, 844 insertions, 0 deletions
diff --git a/dict/dict.go b/dict/dict.go new file mode 100644 index 0000000..f153932 --- /dev/null +++ b/dict/dict.go @@ -0,0 +1,89 @@ +package dict + +import ( + "iter" + "slices" +) + +type Entry[V any] struct { + Key string + Val V +} +type Dict[V any] struct { + s []Entry[V] + m map[string]int +} + +func NewDict[V any]() *Dict[V] { + return &Dict[V]{ + s: make([]Entry[V], 0), + m: make(map[string]int), + } +} + +func (d *Dict[V]) Len() int { + return len(d.s) +} + +func (d *Dict[V]) Set(key string, val V) { + if i, ok := d.m[key]; ok { + d.s[i].Val = val + + return + } + + d.s = append(d.s, Entry[V]{ + Key: key, + Val: val, + }) + + d.m[key] = len(d.s) - 1 +} + +func (d *Dict[V]) Get(key string) (V, bool) { + i, ok := d.m[key] + + if !ok { + var zero V + + return zero, false + } + + return d.s[i].Val, true +} + +func (d *Dict[V]) GetBytes(key []byte) (V, bool) { + if i, ok := d.m[string(key)]; ok { + return d.s[i].Val, true + } + + var zero V + + return zero, false +} + +func (d *Dict[V]) Values() iter.Seq2[int, V] { + return func(yield func(int, V) bool) { + for i, entry := range d.s { + if !yield(i, entry.Val) { + return + } + } + } +} + +func (d *Dict[V]) Sort(cmp func(a, b Entry[V]) int) { + slices.SortFunc(d.s, cmp) + + for i, entry := range d.s { + d.m[entry.Key] = i + } +} + +// func CmpKey[V any](a, b Entry[V]) int { +// return strings.Compare(a.Key, b.Key) +// } + +// func CmpVal[V cmp.Ordered](a, b Entry[V]) int { +// return cmp.Compare(a.Val, b.Val) +// } diff --git a/dict/dict_test.go b/dict/dict_test.go new file mode 100644 index 0000000..c5fb717 --- /dev/null +++ b/dict/dict_test.go @@ -0,0 +1,78 @@ +package dict + +import ( + "fmt" + "slices" + "strings" + "testing" +) + +func TestDict_Sort(t *testing.T) { + d := NewDict[string]() + + d.Set("foo", "foo") + d.Set("bar", "bar") + d.Set("baz", "baz") + + d.Sort(func(a, b Entry[string]) int { + return strings.Compare(a.Key, b.Key) + }) + + expected := []Entry[string]{ + {"bar", "bar"}, + {"baz", "baz"}, + {"foo", "foo"}, + } + + verifyMap(t, d) + + if !slices.Equal(d.s, expected) { + t.Fatalf("expected %v but got %v", expected, d.s) + } +} + +func verifyMap(t *testing.T, dict *Dict[string]) { + if len(dict.m) != len(dict.s) { + t.Fatalf("length missmatch") + } + + for i, s := range dict.s { + v, ok := dict.m[s.Key] + + if !ok { + t.Fatalf("unknown key %s", s) + } + + if v != i { + t.Errorf("expected %d but got %d\n", i, v) + } + } +} + +func BenchmarkDict_Get(b *testing.B) { + d := NewDict[int]() + + for i := 0; i < 1000; i++ { + d.Set(fmt.Sprintf("key_%d", i), i) + } + + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, _ = d.Get("key_500") + } +} + +func BenchmarkDict_GetBytes(b *testing.B) { + d := NewDict[int]() + + for i := 0; i < 1000; i++ { + d.Set(fmt.Sprintf("key_%d", i), i) + } + + keyBytes := []byte("key_500") + + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, _ = d.GetBytes(keyBytes) + } +} diff --git a/dict/go.mod b/dict/go.mod new file mode 100644 index 0000000..23295f8 --- /dev/null +++ b/dict/go.mod @@ -0,0 +1,3 @@ +module go.jknobloc.com/x/dict + +go 1.25 diff --git a/dict/serialize.go b/dict/serialize.go new file mode 100644 index 0000000..da97c06 --- /dev/null +++ b/dict/serialize.go @@ -0,0 +1,55 @@ +package dict + +import ( + "bufio" + "io" + "strings" +) + +func LoadDict[V any](r io.Reader, f func(l string) (string, V)) (*Dict[V], error) { + //var d *Dict[V] + // + //if l, err := linesNotEmpty(r); err != nil { + // return nil, err + //} else { + // d = NewDict[V](l) + //} + + d := NewDict[V]() + + scanner := bufio.NewScanner(r) + + for scanner.Scan() { + line := scanner.Text() + + if err := scanner.Err(); err != nil { + return nil, err + } + + k, v := f(line) + + d.Set(k, v) + } + + return d, nil +} + +func linesNotEmpty(r io.Reader) (int, error) { + scanner := bufio.NewScanner(r) + + count := 0 + + for scanner.Scan() { + line := scanner.Text() + + if strings.TrimSpace(line) != "" { + count++ + } + } + + if err := scanner.Err(); err != nil { + return 0, err + } + + return count, nil +} @@ -2,12 +2,14 @@ go 1.25.0 use ( ./dataset + ./dict ./gpt2 ./llm ./llmc ./mbpe ./onnx ./profile + ./research/entropy ./research/knobloch ./research/lesci ./research/sander diff --git a/research/entropy/cmd/segment/main.go b/research/entropy/cmd/segment/main.go new file mode 100644 index 0000000..cc6aea4 --- /dev/null +++ b/research/entropy/cmd/segment/main.go @@ -0,0 +1,134 @@ +package main + +import ( + "database/sql" + "fmt" + "io" + "log" + "os" + + "github.com/jonasknobloch/mbpe" + + "go.jknobloc.com/x/dict" + "go.jknobloc.com/x/research/entropy" + "go.jknobloc.com/x/shelf" + "go.jknobloc.com/x/tokenizer/bpe" + + _ "github.com/duckdb/duckdb-go/v2" +) + +type NoPreTok struct{} + +func (p *NoPreTok) PreTokenize(s string) []string { + return []string{s} +} + +func main() { + o := shelf.Abs("results/runs-entropy/lesci_minipile_test_256_seed_default/gpt2_256_m000_minipile_seed/lesci.db") + + var db *sql.DB + + if database, err := initDatabase(o); err != nil { + log.Fatal(err) + } else { + db = database + } + + defer db.Close() + + db.SetMaxOpenConns(1) + + var file *os.File + + if f, err := os.Open(shelf.Abs("results/entropy/minipile/dict.txt")); err != nil { + log.Fatal(err) + } else { + file = f + } + + tok := tokenizerMBPE() + + var dict *dict.Dict[*entropy.Entry] + + if d, err := loadDict(file, tok); err != nil { + log.Fatal(err) + } else { + dict = d + } + + if err := entropy.Run(db, dict); err != nil { + log.Fatal(err) + } +} + +func initDatabase(dsn string) (*sql.DB, error) { + var db *sql.DB + + if database, err := sql.Open("duckdb", dsn); err != nil { + return nil, err + } else { + db = database + } + + if err := db.Ping(); err != nil { + _ = db.Close() + + return nil, err + } + + return db, nil +} + +func loadDict(r io.Reader, tok *mbpe.Tokenizer) (*dict.Dict[*entropy.Entry], error) { + return dict.LoadDict[*entropy.Entry](r, func(l string) (string, *entropy.Entry) { + var s string + var n int + + if _, err := fmt.Sscanf(l, "%s %d", &s, &n); err != nil { + panic(err) + } + + ids := tok.Tokenize(s) + + b := to[int, uint8](ids) // TODO that conversion is fucking ass + + return string(b), &entropy.Entry{ + Encoded: s, // TODO encoded or decoded + TokenIDs: b, + LogProbs: make([]float32, len(ids)), + Bounds: make([]bool, len(ids)+1), + N: 0, + } + }) +} + +func to[a bpe.Integer, b bpe.Integer](s []a) []b { + r := make([]b, len(s)) + + // TODO handle overflow + + for i, v := range s { + r[i] = b(v) + } + + return r +} + +func tokenizerMBPE() *mbpe.Tokenizer { + var tok *bpe.Tokenizer + + v := shelf.Abs("tokenizers/minipile/tokenizer_gpt2_256_m000_minipile/vocab.json") + m := shelf.Abs("tokenizers/minipile/tokenizer_gpt2_256_m000_minipile/merges.txt") + + if t, err := bpe.NewTokenizerFromFiles(v, m, bpe.DefaultConfig()); err != nil { + log.Fatal(err) + } else { + tok = t + } + + foo := bpe.MBPE(tok) + + foo.SetPreTokenizer(&NoPreTok{}) + + return foo +} diff --git a/research/entropy/context.go b/research/entropy/context.go new file mode 100644 index 0000000..3993c03 --- /dev/null +++ b/research/entropy/context.go @@ -0,0 +1,53 @@ +package entropy + +import ( + "fmt" + + "go.jknobloc.com/x/dataset" + "go.jknobloc.com/x/llm" +) + +type logProb struct { + document int + token int + value float32 + offset int +} + +func Context(model llm.Causal, tokenizer llm.Tokenizer, data dataset.Reader) error { + evaluatorConfig := llm.EvaluatorConfig{ + BatchSize: 32, + NumWorkers: 16, + } + + tokenBufferConfig := llm.TokenBufferConfig{ + Window: 1024, + Stride: 512, + PadLeft: false, + PadRight: false, + PadTokenID: 256, + } + + eval := llm.NewEvaluator(model, tokenizer, func(job llm.Job, logProbs []float32, tokens []int) []logProb { + r := make([]logProb, len(tokens)) + + for i, token := range tokens { + r[i] = logProb{ + document: job.Document, + token: token, + value: logProbs[i], + offset: job.Position*tokenBufferConfig.Stride + job.Seen + i, + } + } + + return r + }, evaluatorConfig) + + return eval.RunAndCollect("Context", data, tokenBufferConfig, func(r []logProb) error { + for _, l := range r { + fmt.Println(l) // TODO implement + } + + return nil + }) +} diff --git a/research/entropy/entry.go b/research/entropy/entry.go new file mode 100644 index 0000000..3ad2a26 --- /dev/null +++ b/research/entropy/entry.go @@ -0,0 +1,9 @@ +package entropy + +type Entry struct { + Encoded string + TokenIDs []uint8 + LogProbs []float32 + Bounds []bool + N int +} diff --git a/research/entropy/experiment.go b/research/entropy/experiment.go new file mode 100644 index 0000000..b5efad9 --- /dev/null +++ b/research/entropy/experiment.go @@ -0,0 +1,210 @@ +package entropy + +import ( + "cmp" + "database/sql" + "encoding/gob" + "os" + "time" + + "github.com/jonasknobloch/mbpe" + + "go.jknobloc.com/x/dict" + "go.jknobloc.com/x/shelf" + "go.jknobloc.com/x/tui" +) + +type segmentsEntry struct { + Value []string + OK bool +} + +const BufferSize = 1024 + +func Run(db *sql.DB, d *dict.Dict[*Entry]) error { + // m := 0 + // + // for _, v := range d.Values() { + // m = max(m, len(v.TokenIDs)) + // } + // + // if m > BufferSize { + // // TODO handle + // } + + var rows *sql.Rows + + if r, err := db.Query(`SELECT * FROM context`); err != nil { + return err + } else { + rows = r + } + + defer rows.Close() + + bufferTokens := make([]uint8, 0, BufferSize) + bufferLogProbs := make([]float32, 0, BufferSize) + + doc := -1 + + maxTokens := 100000 + currentTokens := 0 + + pb := tui.NewProgressBar("Adding LogProbs", 20, 58508191, time.Now()) + + pb.Start(1*time.Second, func() int { + return pb.Completed() + }) + + for rows.Next() { + if currentTokens == maxTokens { + break + } + + currentTokens++ + + var uid int + var token int + var logProb float32 + var pos int + + if err := rows.Scan(&uid, &token, &logProb, &pos); err != nil { + return err + } + + if uid != doc || len(bufferTokens) == BufferSize { + bufferTokens = bufferTokens[:0] + bufferLogProbs = bufferLogProbs[:0] + + doc = uid + } + + bufferTokens = append(bufferTokens, uint8(token)) + bufferLogProbs = append(bufferLogProbs, logProb) + + n := len(bufferTokens) + + // fmt.Print(n) + + for i := range n { + entry, ok := d.GetBytes(bufferTokens[i:n]) + + if !ok { + continue + } + + for j := range entry.LogProbs { + entry.LogProbs[j] += bufferLogProbs[j] + } + + entry.N++ + } + + pb.Add(1) + } + + if err := rows.Err(); err != nil { + return err + } + + // debug + d.Sort(func(a, b dict.Entry[*Entry]) int { + return cmp.Compare(b.Val.N, a.Val.N) + }) + + pb.Close() + + for _, e := range d.Values() { + spikes(e) + } + + if err := serializeSegments(d); err != nil { + return err + } + + return nil +} + +func spikes(e *Entry) { + for i := range len(e.LogProbs) { + if i == 0 { + continue + } + + if i == 1 && e.TokenIDs[0] == 220 { + continue + } + + // TODO threshold + + if d := e.LogProbs[i-1] - e.LogProbs[i]; d > 0 { + e.Bounds[i] = true + } + } +} + +func serializeSegments(d *dict.Dict[*Entry]) error { + r := make(map[string]segmentsEntry, d.Len()) + + for _, e := range d.Values() { + segments := make([]string, 0) + + decoded := decode(e.Encoded) + + i := 0 + + for j, b := range e.Bounds { + if b { + segments = append(segments, encode(decoded[i:j])) + + i = j + } + } + + segments = append(segments, encode(decoded[i:])) // TODO or just use OK false? + + // debug + if i == 0 { + continue + } + + r[e.Encoded] = segmentsEntry{ + Value: segments, + OK: true, + } + } + + var file *os.File + + if f, err := os.Create(shelf.Abs("results/entropy/minipile/segments.gob")); err != nil { + return err + } else { + file = f + } + + defer file.Close() + + enc := gob.NewEncoder(file) + + return enc.Encode(r) +} + +func encode(s string) string { + result := "" + + for _, b := range []byte(s) { + result += mbpe.BytesChar[b] + } + + return result +} + +func decode(s string) string { + result := make([]byte, 0) + + for _, r := range s { + result = append(result, mbpe.CharBytes[string(r)]) + } + + return string(result) +} diff --git a/research/entropy/go.mod b/research/entropy/go.mod new file mode 100644 index 0000000..ab7ba3d --- /dev/null +++ b/research/entropy/go.mod @@ -0,0 +1,32 @@ +module go.jknobloc.com/x/research/entropy + +go 1.25.0 + +require github.com/duckdb/duckdb-go/v2 v2.10501.0 + +require ( + github.com/apache/arrow-go/v18 v18.5.1 // indirect + github.com/duckdb/duckdb-go-bindings v0.10501.0 // indirect + github.com/duckdb/duckdb-go-bindings/lib/darwin-amd64 v0.10501.0 // indirect + github.com/duckdb/duckdb-go-bindings/lib/darwin-arm64 v0.10501.0 // indirect + github.com/duckdb/duckdb-go-bindings/lib/linux-amd64 v0.10501.0 // indirect + github.com/duckdb/duckdb-go-bindings/lib/linux-arm64 v0.10501.0 // indirect + github.com/duckdb/duckdb-go-bindings/lib/windows-amd64 v0.10501.0 // indirect + github.com/go-viper/mapstructure/v2 v2.5.0 // indirect + github.com/goccy/go-json v0.10.5 // indirect + github.com/google/flatbuffers v25.12.19+incompatible // indirect + github.com/google/go-cmp v0.7.0 // indirect + github.com/google/uuid v1.6.0 // indirect + github.com/klauspost/compress v1.18.3 // indirect + github.com/klauspost/cpuid/v2 v2.3.0 // indirect + github.com/pierrec/lz4/v4 v4.1.25 // indirect + github.com/zeebo/xxh3 v1.1.0 // indirect + golang.org/x/exp v0.0.0-20260112195511-716be5621a96 // indirect + golang.org/x/mod v0.33.0 // indirect + golang.org/x/sync v0.20.0 // indirect + golang.org/x/sys v0.41.0 // indirect + golang.org/x/telemetry v0.0.0-20260209163413-e7419c687ee4 // indirect + golang.org/x/tools v0.42.0 // indirect + golang.org/x/xerrors v0.0.0-20240903120638-7835f813f4da // indirect + gonum.org/v1/gonum v0.17.0 // indirect +) diff --git a/research/entropy/go.sum b/research/entropy/go.sum new file mode 100644 index 0000000..70ff7c1 --- /dev/null +++ b/research/entropy/go.sum @@ -0,0 +1,36 @@ +github.com/andybalholm/brotli v1.2.0 h1:ukwgCxwYrmACq68yiUqwIWnGY0cTPox/M94sVwToPjQ= +github.com/apache/arrow-go/v18 v18.5.1 h1:yaQ6zxMGgf9YCYw4/oaeOU3AULySDlAYDOcnr4LdHdI= +github.com/apache/thrift v0.22.0 h1:r7mTJdj51TMDe6RtcmNdQxgn9XcyfGDOzegMDRg47uc= +github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM= +github.com/duckdb/duckdb-go-bindings v0.10501.0 h1:BR21HkcALr9Lm+Ios2vEPaaB5oRRxGJHONzkS0bnOKE= +github.com/duckdb/duckdb-go-bindings/lib/darwin-amd64 v0.10501.0 h1:InnDiz/iBHUzwI/4xkigTq6PRrIx+9L+eC2NfCShgWc= +github.com/duckdb/duckdb-go-bindings/lib/darwin-arm64 v0.10501.0 h1:XLMUi/9QJcN8Bp77ML/QPwynX8f9RAg4VUiTdPzRUEU= +github.com/duckdb/duckdb-go-bindings/lib/linux-amd64 v0.10501.0 h1:td84w8XucSPQoxGC84RYIxTu1+RV+fjeFIHVk3GLXog= +github.com/duckdb/duckdb-go-bindings/lib/linux-arm64 v0.10501.0 h1:XLw03uWhdQvAFU6unJ2MtQvSFi4LNCfhiOM644MyAHw= +github.com/duckdb/duckdb-go-bindings/lib/windows-amd64 v0.10501.0 h1:jhhOonew2VOcTN4f+BOlnOTUgHEp797Uee2Tq8xMno8= +github.com/duckdb/duckdb-go/v2 v2.10501.0 h1:vYgvKBfotrZqpBESqHXYF5NVlbaYHRf0VrQEXtb/jnU= +github.com/go-viper/mapstructure/v2 v2.5.0 h1:vM5IJoUAy3d7zRSVtIwQgBj7BiWtMPfmPEgAXnvj1Ro= +github.com/goccy/go-json v0.10.5 h1:Fq85nIqj+gXn/S5ahsiTlK3TmC85qgirsdTP/+DeaC4= +github.com/golang/snappy v1.0.0 h1:Oy607GVXHs7RtbggtPBnr2RmDArIsAefDwvrdWvRhGs= +github.com/google/flatbuffers v25.12.19+incompatible h1:haMV2JRRJCe1998HeW/p0X9UaMTK6SDo0ffLn2+DbLs= +github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/klauspost/asmfmt v1.3.2 h1:4Ri7ox3EwapiOjCki+hw14RyKk201CN4rzyCJRFLpK4= +github.com/klauspost/compress v1.18.3 h1:9PJRvfbmTabkOX8moIpXPbMMbYN60bWImDDU7L+/6zw= +github.com/klauspost/cpuid/v2 v2.3.0 h1:S4CRMLnYUhGeDFDqkGriYKdfoFlDnMtqTiI/sFzhA9Y= +github.com/minio/asm2plan9s v0.0.0-20200509001527-cdd76441f9d8 h1:AMFGa4R4MiIpspGNG7Z948v4n35fFGB3RR3G/ry4FWs= +github.com/minio/c2goasm v0.0.0-20190812172519-36a3d3bbc4f3 h1:+n/aFZefKZp7spd8DFdX7uMikMLXX4oubIzJF4kv/wI= +github.com/pierrec/lz4/v4 v4.1.25 h1:kocOqRffaIbU5djlIBr7Wh+cx82C0vtFb0fOurZHqD0= +github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 h1:Jamvg5psRIccs7FGNTlIRMkT8wgtp5eCXdBlqhYGL6U= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/zeebo/assert v1.3.0 h1:g7C04CbJuIDKNPFHmsk4hwZDO5O+kntRxzaUoNXj+IQ= +github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs= +golang.org/x/exp v0.0.0-20260112195511-716be5621a96 h1:Z/6YuSHTLOHfNFdb8zVZomZr7cqNgTJvA8+Qz75D8gU= +golang.org/x/mod v0.33.0 h1:tHFzIWbBifEmbwtGz65eaWyGiGZatSrT9prnU8DbVL8= +golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= +golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k= +golang.org/x/telemetry v0.0.0-20260209163413-e7419c687ee4 h1:bTLqdHv7xrGlFbvf5/TXNxy/iUwwdkjhqQTJDjW7aj0= +golang.org/x/tools v0.42.0 h1:uNgphsn75Tdz5Ji2q36v/nsFSfR/9BRFvqhGBaJGd5k= +golang.org/x/xerrors v0.0.0-20240903120638-7835f813f4da h1:noIWHXmPHxILtqtCOPIhSt0ABwskkZKjD3bXGnZGpNY= +gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= diff --git a/tokenizer/byte/tokenizer.go b/tokenizer/byte/tokenizer.go new file mode 100644 index 0000000..98982d0 --- /dev/null +++ b/tokenizer/byte/tokenizer.go @@ -0,0 +1,94 @@ +package byte + +import ( + "encoding/json" + "io" + + "github.com/jonasknobloch/mbpe" +) + +// type parameter for ID type in bpe tokenzier ?! + +// Name? Byte is bad; it is a bijection but so are most if not all tokenizers; its the Alphabet? +// It would be compativle with GenericTokenizer[uint8] without pretokenization + +// Tokenizer is a byte level tokenzier covering the all 2^8 bytes; Essentially each individual byte is mapped to a token ID. +// Serialization follows the standard vocab.json (with HF byte replacements) while omitting merges.txt +// Pre-tokenization is not necessary; however the tokenized strings can be much larger -> check allocations +type Tokenizer struct { + atoi map[byte]uint8 + itoa map[uint8]byte +} + +func (t *Tokenizer) Encode(s string) []byte { + b := make([]byte, len(s)) + + for i := range len(s) { + v, ok := t.atoi[s[i]] + + if !ok { + panic("unknown byte") + } + + b[i] = v + } + + return b +} + +func (t *Tokenizer) Decode(ids []uint8) string { + b := make([]byte, len(ids)) + + for i, id := range ids { + v, ok := t.itoa[id] + + if !ok { + panic("unknown token ID") + } + + b[i] = v + } + + return string(b) +} + +func (t *Tokenizer) Tokenize(s string) []int { + ids := t.Encode(s) + + r := make([]int, len(ids)) + + for i, id := range ids { + r[i] = int(id) + } + + return r +} + +func NewTokenizer(vocab io.Reader) (*Tokenizer, error) { + v := make(map[string]uint8) + + decoder := json.NewDecoder(vocab) + + if err := decoder.Decode(&v); err != nil { + return nil, err + } + + if len(v) != 256 { + panic("vocabulary size != 256") + } + + atoi := make(map[byte]uint8, len(v)) + itoa := make(map[uint8]byte, len(v)) + + for char, id := range v { + b := mbpe.CharBytes[char] + + atoi[b] = id + itoa[id] = b + } + + return &Tokenizer{ + atoi: atoi, + itoa: itoa, + }, nil +} diff --git a/tokenizer/byte/tokenizer_test.go b/tokenizer/byte/tokenizer_test.go new file mode 100644 index 0000000..d2c8c6c --- /dev/null +++ b/tokenizer/byte/tokenizer_test.go @@ -0,0 +1,49 @@ +package byte + +import ( + "fmt" + "os" + "testing" + + "go.jknobloc.com/x/shelf" +) + +func TestNewTokenizer_Encode(t *testing.T) { + var file *os.File + + if f, err := os.Open(shelf.Abs("tokenizers/minipile/tokenizer_gpt2_256_m000_minipile/vocab.json")); err != nil { + t.Fatal(err) + } else { + file = f + } + + defer file.Close() + + tok, err := NewTokenizer(file) + + if err != nil { + t.Fatal(err) + } + + fmt.Println(tok.Encode(" \nabc")) +} + +func TestNewTokenizer_Decode(t *testing.T) { + var file *os.File + + if f, err := os.Open(shelf.Abs("tokenizers/minipile/tokenizer_gpt2_256_m000_minipile/vocab.json")); err != nil { + t.Fatal(err) + } else { + file = f + } + + defer file.Close() + + tok, err := NewTokenizer(file) + + if err != nil { + t.Fatal(err) + } + + fmt.Println(tok.Decode([]uint8{220, 198, 64, 65, 66})) +} |
