From 8fd3f8aeec85f47d43ef36a81456d48640ca33f9 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Wed, 9 Sep 2026 02:40:41 +0200 Subject: Add dict module --- dict/dict.go | 85 +++++++++++++++++++++++++++++++++++++++++++++++++++++++ dict/dict_test.go | 78 ++++++++++++++++++++++++++++++++++++++++++++++++++ dict/go.mod | 3 ++ go.work | 1 + 4 files changed, 167 insertions(+) create mode 100644 dict/dict.go create mode 100644 dict/dict_test.go create mode 100644 dict/go.mod diff --git a/dict/dict.go b/dict/dict.go new file mode 100644 index 0000000..bc0f8bd --- /dev/null +++ b/dict/dict.go @@ -0,0 +1,85 @@ +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]) 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/go.work b/go.work index 9cdd338..b4654b7 100644 --- a/go.work +++ b/go.work @@ -2,6 +2,7 @@ go 1.25.0 use ( ./dataset + ./dict ./gpt2 ./llm ./llmc -- cgit v1.2.3 From 3c5f4c4279e6c9d6479e40eff4fbe1add59c1b3a Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Wed, 9 Sep 2026 02:41:24 +0200 Subject: WIP --- go.work | 1 + research/entropy/cmd/dict3/main.go | 159 +++++++++++++++++ research/entropy/cmd/main.go | 221 +++++++++++++++++++++++ research/entropy/context.go | 53 ++++++ research/entropy/entry.go | 8 + research/entropy/experiment.go | 119 +++++++++++++ research/entropy/go.mod | 32 ++++ research/entropy/go.sum | 36 ++++ research/entropy/segment.go | 346 +++++++++++++++++++++++++++++++++++++ research/entropy/segmenter.go | 94 ++++++++++ tokenizer/bpe/utility.go | 4 + 11 files changed, 1073 insertions(+) create mode 100644 research/entropy/cmd/dict3/main.go create mode 100644 research/entropy/cmd/main.go create mode 100644 research/entropy/context.go create mode 100644 research/entropy/entry.go create mode 100644 research/entropy/experiment.go create mode 100644 research/entropy/go.mod create mode 100644 research/entropy/go.sum create mode 100644 research/entropy/segment.go create mode 100644 research/entropy/segmenter.go diff --git a/go.work b/go.work index b4654b7..2ec50a9 100644 --- a/go.work +++ b/go.work @@ -9,6 +9,7 @@ use ( ./mbpe ./onnx ./profile + ./research/entropy ./research/knobloch ./research/lesci ./research/sander diff --git a/research/entropy/cmd/dict3/main.go b/research/entropy/cmd/dict3/main.go new file mode 100644 index 0000000..a6b0211 --- /dev/null +++ b/research/entropy/cmd/dict3/main.go @@ -0,0 +1,159 @@ +package main + +import ( + "bufio" + "database/sql" + "fmt" + "io" + "log" + "os" + + _ "github.com/duckdb/duckdb-go/v2" + "github.com/jonasknobloch/mbpe" + "go.jknobloc.com/x/dict" + "go.jknobloc.com/x/research/entropy" + "go.jknobloc.com/x/tokenizer/bpe" + + "go.jknobloc.com/x/shelf" +) + +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/knobloch/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) { + d := dict.NewDict[*entropy.Entry]() + + // TODO we could count non empty lines to pre alloc dict + + scanner := bufio.NewScanner(r) + + for scanner.Scan() { + line := scanner.Text() + + if err := scanner.Err(); err != nil { + return nil, err + } + + var s string + var n int + + if _, err := fmt.Sscanf(line, "%s %d", &s, &n); err != nil { + return nil, err + } + + // TODO see cmd/main.go we currently need mbpe tok ?? + + if n < 100 { + continue + } + + ids := tok.Tokenize(s) + + if len(ids) < 4 { + continue + } + + b := to[int, uint8](ids) // TODO that conversion is fucking ass + + // TODO vs decode each buffer suffix and use s as map key + + d.Set(string(b), &entropy.Entry{ + Encoded: s, // TODO encoded or decoded + TokenIDs: b, + LogProbs: make([]float32, len(ids)), + N: 0, + }) + } + + return d, nil +} + +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/cmd/main.go b/research/entropy/cmd/main.go new file mode 100644 index 0000000..cba8c33 --- /dev/null +++ b/research/entropy/cmd/main.go @@ -0,0 +1,221 @@ +package main + +import ( + "database/sql" + "fmt" + "log" + + _ "github.com/duckdb/duckdb-go/v2" + "go.jknobloc.com/x/research/entropy" + + "github.com/jonasknobloch/mbpe" + "go.jknobloc.com/x/dataset" + "go.jknobloc.com/x/gpt2" + "go.jknobloc.com/x/shelf" + "go.jknobloc.com/x/tokenizer/bpe" +) + +type NoPreTok struct{} + +func (p *NoPreTok) PreTokenize(s string) []string { + return []string{s} +} + +const CorpusMean = -0.7749100363022372 + +const Threshold = -2.0 + +func main() { + dict := mbpe.NewDict() + + if err := dict.Load(shelf.Abs("results/knobloch/minipile/dict.txt")); err != nil { + log.Fatal(err) + } + + if err := gpt2.InitializeEnvironment(); err != nil { + log.Fatal(err) + } + + // m := model() + // t := tokenizer() + // d := data() + // + // if err := entropy.Context(m, t, d); err != nil { + // log.Fatal(err) + // } + + 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) + + foo := tokenizerMBPE() + + type candidate struct { + src string + ids []int + means []float64 + occurrences int64 + } + + candidates := make([]candidate, 0) + baseline := entropy.NewBaseline() + + for i, c := range dict.Items()[10000:] { + if i > 10 { + break + } + + s := c.Src() + + if len(s) < 10 { + continue + } + + ids := foo.Tokenize(s) + + means, occurrences, err := entropy.MeanLogProbs(ids, db) + + if err != nil { + log.Fatal(err) + } + + if occurrences == 0 { + continue + } + + { + occurrences, err := entropy.LogProbs(ids, db) + + if err != nil { + log.Fatal(err) + } + + fmt.Println(ids) + fmt.Println(occurrences) + } + + candidates = append(candidates, candidate{s, ids, means, occurrences}) + + baseline.Add(means) + } + + for _, c := range candidates { + fmt.Println(c.ids) + fmt.Println(foo.Decoder().Decode([]string{c.src})) + fmt.Println(c.means, c.occurrences) + + fmt.Println("excess ", entropy.Excess(c.means, Threshold)) + fmt.Println("spikes ", entropy.Spikes(c.means)) + fmt.Println("deviations", baseline.Deviations(c.means)) + + fmt.Println() + fmt.Println() + fmt.Println() + } + + if err := gpt2.DestroyEnvironment(); err != nil { + log.Fatal(err) + } +} + +// func must[T any](v T, err error) T { +// if err != nil { +// log.Fatal(err) +// } +// +// return v +// } + +func model() *gpt2.Model { + cfg := gpt2.ConfigDefault() + + cfg.VocabSize = 256 + + opts := gpt2.Options{ + WithCache: false, + WithLogits: false, + WithLogProbs: true, + } + + m := gpt2.NewModel(shelf.Abs("models/mbpe/minipile-byte/gpt2_256_m000_minipile_seed/model_eval.onnx"), cfg, opts) + + if err := m.Init(); err != nil { + log.Fatal(err) + } + + return m +} + +func tokenizer() *bpe.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 + } + + return tok +} + +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 +} + +func data() dataset.Reader { + var miniPile *dataset.ParquetReader + + if r, err := dataset.NewParquetReader(shelf.Abs("data/minipile/validation")); err != nil { + log.Fatal(err) + } else { + miniPile = r + } + + return dataset.NewClampedReader(miniPile, 10) +} + +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 +} 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..746e34f --- /dev/null +++ b/research/entropy/entry.go @@ -0,0 +1,8 @@ +package entropy + +type Entry struct { + Encoded string + TokenIDs []uint8 + LogProbs []float32 + N int +} diff --git a/research/entropy/experiment.go b/research/entropy/experiment.go new file mode 100644 index 0000000..1e4814c --- /dev/null +++ b/research/entropy/experiment.go @@ -0,0 +1,119 @@ +package entropy + +import ( + "cmp" + "database/sql" + + "go.jknobloc.com/x/dict" +) + +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 + // } + + rows, err := db.Query(`SELECT * FROM context`) + + if err != nil { + return err + } + + defer rows.Close() + + buffer := make([]uint8, 0, BufferSize) + bufferLogProbs := make([]float32, 0, BufferSize) + + doc := -1 + + maxTokens := 10000000 + currentToken := 0 + + for rows.Next() { + if currentToken == maxTokens { + // break + } + + currentToken++ + + 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(buffer) == BufferSize { + buffer = buffer[:0] + bufferLogProbs = bufferLogProbs[:0] + + doc = uid + } + + buffer = append(buffer, uint8(token)) + bufferLogProbs = append(bufferLogProbs, logProb) + + n := len(buffer) + + // fmt.Print(n) + + for i := range n { + entry, ok := d.GetBytes(buffer[i:n]) + + if !ok { + continue + } + + for j := range entry.LogProbs { + entry.LogProbs[j] += bufferLogProbs[j] + } + + entry.N++ + } + } + + 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) + }) + + // for _, e := range d.Values() { + // if e.N < 100 { + // break + // } + // + // fmt.Println(e.Encoded) + // fmt.Println(spikes(e.LogProbs)) + // } + + return nil +} + +func spikes(logProbs []float32) []bool { + b := make([]bool, len(logProbs)) + + for i := range len(logProbs) { + if i == 0 { + continue + } + + if d := logProbs[i-1] - logProbs[i]; d > 0 { + b[i] = true // TODO threshold + } + } + + return b +} 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/research/entropy/segment.go b/research/entropy/segment.go new file mode 100644 index 0000000..bdab576 --- /dev/null +++ b/research/entropy/segment.go @@ -0,0 +1,346 @@ +package entropy + +import ( + "database/sql" + "math" +) + +// CREATE TABLE context(uid INTEGER, token INTEGER, logprob FLOAT, pos INTEGER) +// +// sqlMeanLogProbs matches an exact token id sequence of any length, bound as a +// single list parameter, and returns one row per position: the position, the +// mean logprob over all occurrences, and the occurrence count. +// +// seq unrolls the parameter into (position, token) pairs. Joining context on +// token alone gives every position that matches some element of the sequence; +// pos - i + 1 is the start a match would have to begin at for that element to +// line up, so counting rows per (uid, start) says how many elements agree +// there. A start where all of them agree is an occurrence. +const sqlMeanLogProbs = ` +WITH seq AS ( + SELECT i, s[i] AS token + FROM (SELECT ?::INTEGER[] AS s) t, range(1, len(s) + 1) r(i) +), +cand AS ( + SELECT s.i AS i, + c.logprob AS logprob, + count(*) OVER (PARTITION BY c.uid, c.pos - s.i + 1) AS matched + FROM context c + JOIN seq s ON c.token = s.token +) +SELECT i, avg(logprob) AS mean, count(*) AS occurrences +FROM cand +WHERE matched = (SELECT count(*) FROM seq) +GROUP BY i +ORDER BY i` + +// sqlLogProbs matches the same sequences as sqlMeanLogProbs but skips the +// aggregation, returning one row per position per occurrence. +const sqlLogProbs = ` +WITH seq AS ( + SELECT i, s[i] AS token + FROM (SELECT ?::INTEGER[] AS s) t, range(1, len(s) + 1) r(i) +), +cand AS ( + SELECT c.uid AS uid, + c.pos - s.i + 1 AS start, + s.i AS i, + c.logprob AS logprob, + count(*) OVER (PARTITION BY c.uid, c.pos - s.i + 1) AS matched + FROM context c + JOIN seq s ON c.token = s.token +) +SELECT uid, start, i, logprob +FROM cand +WHERE matched = (SELECT count(*) FROM seq) +ORDER BY uid, start, i` + +// Occurrence is one match of a sequence, holding the logprob of every position +// in the context it was found in. +type Occurrence struct { + Document int + Offset int + Values []float64 +} + +// LogProbs returns every occurrence of the exact id sequence separately, +// leaving the contexts uncombined. +func LogProbs(ids []int, db *sql.DB) ([]Occurrence, error) { + seq := make([]int32, len(ids)) // the driver panics on []int + + for i, id := range ids { + seq[i] = int32(id) + } + + rows, err := db.Query(sqlLogProbs, seq) + + if err != nil { + return nil, err + } + + defer rows.Close() + + r := make([]Occurrence, 0) + + for rows.Next() { + var document, offset, pos int + var value float64 + + if err := rows.Scan(&document, &offset, &pos, &value); err != nil { + return nil, err + } + + if pos == 1 { + r = append(r, Occurrence{ + Document: document, + Offset: offset, + Values: make([]float64, 0, len(ids)), + }) + } + + if len(r) == 0 { + continue + } + + last := &r[len(r)-1] + + last.Values = append(last.Values, value) + } + + if err := rows.Err(); err != nil { + return nil, err + } + + return r, nil +} + +// MeanLogProbs returns the mean logprob per position over all occurrences of +// the exact id sequence, together with the number of occurrences. Occurrences +// of zero yield a nil slice: the sequence was never seen. +func MeanLogProbs(ids []int, db *sql.DB) ([]float64, int64, error) { + seq := make([]int32, len(ids)) // the driver panics on []int + + for i, id := range ids { + seq[i] = int32(id) + } + + rows, err := db.Query(sqlMeanLogProbs, seq) + + if err != nil { + return nil, 0, err + } + + defer rows.Close() + + values := make([]float64, 0, len(ids)) + + var occurrences int64 + + for rows.Next() { + var pos int + var mean float64 + var n int64 + + if err := rows.Scan(&pos, &mean, &n); err != nil { + return nil, 0, err + } + + values = append(values, mean) + + occurrences = n + } + + if err := rows.Err(); err != nil { + return nil, 0, err + } + + if len(values) == 0 { + return nil, 0, nil + } + + return values, occurrences, nil +} + +// Excess weights each position by how far it falls below threshold, in nats. +// Fires wherever surprisal is high in absolute terms, which in practice means +// the uncertain zone at the start of a chunk. +func Excess(values []float64, threshold float64) []float64 { + r := make([]float64, len(values)) + + for i, v := range values { + if d := threshold - v; d > 0 { + r[i] = d + } + } + + return r +} + +// Spikes weights each position by how much more surprising it is than the one +// before it, in nats. Catches boundaries an absolute threshold misses, since +// surprisal decays inside a unit and a boundary shows up as uncertainty rising +// again, however low the preceding tail got. Position 0 has no predecessor and +// is always 0. +func Spikes(values []float64) []float64 { + r := make([]float64, len(values)) + + for i := 1; i < len(values); i++ { + if d := values[i-1] - values[i]; d > 0 { + r[i] = d + } + } + + return r +} + +// Baseline accumulates mean logprobs per position index across sequences. +// Surprisal decays sharply with position -- onsets are uncertain, tails are +// forced -- so the raw value at a position says little on its own. Deviations +// measures how far a sequence departs from what its positions usually look +// like, which also cancels the downward pull the exact-sequence lookup puts on +// interior positions. +type Baseline struct { + sum []float64 + sumSq []float64 + noise []float64 + counts []int +} + +// MinVarianceShare floors the noise correction: however much sampling noise is +// estimated, a position keeps at least this share of its observed variance. +// Without it a position where every candidate agrees would correct to near-zero +// variance and produce enormous deviations from nothing. +const MinVarianceShare = 0.1 + +func NewBaseline() *Baseline { + return &Baseline{} +} + +func (b *Baseline) grow(n int) { + for len(b.counts) < n { + b.sum = append(b.sum, 0) + b.sumSq = append(b.sumSq, 0) + b.noise = append(b.noise, 0) + b.counts = append(b.counts, 0) + } +} + +func (b *Baseline) Add(values []float64) { + b.grow(len(values)) + + for i, v := range values { + b.sum[i] += v + b.sumSq[i] += v * v + b.counts[i]++ + } +} + +// AddCandidate is Add for a candidate whose means were estimated from n +// occurrences with the given per-position variances across those occurrences. +// The spread of candidate means overstates how much candidates really differ, +// by the sampling variance of the means themselves; recording it lets stats +// subtract it instead of passing the inflation on to every deviation. +func (b *Baseline) AddCandidate(values, variances []float64, n int) { + b.Add(values) + + if n < 2 { + return + } + + for i, v := range variances { + if i >= len(b.noise) { + break + } + + b.noise[i] += v / float64(n) + } +} + +func (b *Baseline) stats(pos int) (float64, float64, bool) { + if pos >= len(b.counts) || b.counts[pos] < 2 { + return 0, 0, false + } + + n := float64(b.counts[pos]) + mean := b.sum[pos] / n + + variance := (b.sumSq[pos] - n*mean*mean) / (n - 1) + + if variance <= 0 { + return mean, 0, false + } + + if corrected := variance - b.noise[pos]/n; corrected > variance*MinVarianceShare { + variance = corrected + } else { + variance = variance * MinVarianceShare + } + + return mean, math.Sqrt(variance), true +} + +// Summarize reduces the occurrences of one candidate to per-position means and +// variances, the inputs AddCandidate needs. +func Summarize(occurrences []Occurrence) ([]float64, []float64, int) { + if len(occurrences) == 0 { + return nil, nil, 0 + } + + n := len(occurrences[0].Values) + + means := make([]float64, n) + variances := make([]float64, n) + + for _, o := range occurrences { + for i, v := range o.Values { + if i >= n { + break + } + + means[i] += v + variances[i] += v * v + } + } + + for i := range means { + sum := means[i] + sumSq := variances[i] + + means[i] = sum / float64(len(occurrences)) + + if len(occurrences) < 2 { + variances[i] = 0 + + continue + } + + variances[i] = (sumSq - float64(len(occurrences))*means[i]*means[i]) / float64(len(occurrences)-1) + + if variances[i] < 0 { + variances[i] = 0 + } + } + + return means, variances, len(occurrences) +} + +// Deviations scores each position by how many standard deviations more +// surprising it is than the same position usually is. Positions at or below +// their positional mean score 0, as they carry no evidence of a boundary. +func (b *Baseline) Deviations(values []float64) []float64 { + r := make([]float64, len(values)) + + for i, v := range values { + mean, std, ok := b.stats(i) + + if !ok { + continue + } + + if d := (mean - v) / std; d > 0 { + r[i] = d + } + } + + return r +} diff --git a/research/entropy/segmenter.go b/research/entropy/segmenter.go new file mode 100644 index 0000000..1a15111 --- /dev/null +++ b/research/entropy/segmenter.go @@ -0,0 +1,94 @@ +package entropy + +import ( + "database/sql" +) + +type Tokenizer interface { + Tokenize(string) []int +} + +// Criterion turns per-position mean logprobs into per-position boundary +// weights, zero wherever there is no evidence of a boundary. +type Criterion func([]float64) []float64 + +// ExcessCriterion is Excess bound to a threshold. Weights come out in nats +// below that threshold, so a Segmenter using it wants a cutoff of 0. +func ExcessCriterion(threshold float64) Criterion { + return func(values []float64) []float64 { + return Excess(values, threshold) + } +} + +// Segmenter turns surprisal into segments, satisfying the mbpe Segmenter +// interface so it can stand in for Morfessor without touching the trainer. +// +// A boundary is placed wherever the criterion weighs a position above cutoff. +// Note that cutoff applies to the weights, not to the logprobs: with Spikes it +// is a minimum rise in nats and belongs above zero, with ExcessCriterion the +// threshold is already baked in and cutoff belongs at zero. +type Segmenter struct { + tokenizer Tokenizer + db *sql.DB + criterion Criterion + cutoff float64 +} + +func NewSegmenter(tokenizer Tokenizer, db *sql.DB, criterion Criterion, cutoff float64) *Segmenter { + return &Segmenter{ + tokenizer: tokenizer, + db: db, + criterion: criterion, + cutoff: cutoff, + } +} + +// Segment reports the segmentation of compound and whether surprisal was known +// for it at all. Compounds the model never saw come back unsegmented and not +// ok, which the trainer already treats as a reason to drop alpha to zero. +// +// The leading whitespace marker the trainer strips is put back before the +// lookup: the logprobs were measured on running text, so the first character +// of a word is only predictable given the space in front of it. +func (s *Segmenter) Segment(compound string) ([]string, bool) { + runes := []rune(compound) + + if len(runes) < 2 { + return []string{compound}, false + } + + ids := s.tokenizer.Tokenize("Ġ" + compound) + + if len(ids) != len(runes)+1 { + return []string{compound}, false // not one token per rune, cannot align + } + + values, occurrences, err := MeanLogProbs(ids, s.db) + + if err != nil || occurrences == 0 || len(values) != len(ids) { + return []string{compound}, false + } + + weights := s.criterion(values) + + segments := make([]string, 0, 4) + + start := 0 + + // ids[0] is the whitespace marker and ids[1] the first rune of compound, so + // a boundary there is the start of the chunk and carries no information. + // Skipping it also drops the onset spike, which dominates every sequence. + for i := 2; i < len(ids); i++ { + if weights[i] <= s.cutoff { + continue + } + + segments = append(segments, string(runes[start:i-1])) + + start = i - 1 + } + + segments = append(segments, string(runes[start:])) + + return segments, true +} diff --git a/tokenizer/bpe/utility.go b/tokenizer/bpe/utility.go index f8fcdbc..79e6faf 100644 --- a/tokenizer/bpe/utility.go +++ b/tokenizer/bpe/utility.go @@ -19,6 +19,10 @@ func NewTokenizerFromFiles(vocab, merges string, cfg Config) (*Tokenizer, error) return NewTokenizer(tokenizer, cfg), nil } +func MBPE(t *Tokenizer) *mbpe.Tokenizer { + return t.mbpe +} + func Vocab(t *Tokenizer) []string { m, ok := t.mbpe.Model().(*mbpe.MBPE) -- cgit v1.2.3 From 123a37a04b0076c569cbed7c2669ab1087e2cba7 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Wed, 9 Sep 2026 02:43:55 +0200 Subject: WIP --- tokenizer/byte/tokenizer.go | 94 ++++++++++++++++++++++++++++++++++++++++ tokenizer/byte/tokenizer_test.go | 49 +++++++++++++++++++++ 2 files changed, 143 insertions(+) create mode 100644 tokenizer/byte/tokenizer.go create mode 100644 tokenizer/byte/tokenizer_test.go 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})) +} -- cgit v1.2.3 From f998fee7427403698f9262d1e92d1a307c652094 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Thu, 10 Sep 2026 01:59:34 +0200 Subject: WIP --- dict/dict.go | 4 + dict/serialize.go | 55 ++++++ research/entropy/cmd/dict3/main.go | 159 ---------------- research/entropy/cmd/main.go | 221 ---------------------- research/entropy/cmd/segment/main.go | 134 ++++++++++++++ research/entropy/entry.go | 1 + research/entropy/experiment.go | 145 ++++++++++++--- research/entropy/segment.go | 346 ----------------------------------- research/entropy/segmenter.go | 94 ---------- 9 files changed, 312 insertions(+), 847 deletions(-) create mode 100644 dict/serialize.go delete mode 100644 research/entropy/cmd/dict3/main.go delete mode 100644 research/entropy/cmd/main.go create mode 100644 research/entropy/cmd/segment/main.go delete mode 100644 research/entropy/segment.go delete mode 100644 research/entropy/segmenter.go diff --git a/dict/dict.go b/dict/dict.go index bc0f8bd..f153932 100644 --- a/dict/dict.go +++ b/dict/dict.go @@ -21,6 +21,10 @@ func NewDict[V any]() *Dict[V] { } } +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 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 +} diff --git a/research/entropy/cmd/dict3/main.go b/research/entropy/cmd/dict3/main.go deleted file mode 100644 index a6b0211..0000000 --- a/research/entropy/cmd/dict3/main.go +++ /dev/null @@ -1,159 +0,0 @@ -package main - -import ( - "bufio" - "database/sql" - "fmt" - "io" - "log" - "os" - - _ "github.com/duckdb/duckdb-go/v2" - "github.com/jonasknobloch/mbpe" - "go.jknobloc.com/x/dict" - "go.jknobloc.com/x/research/entropy" - "go.jknobloc.com/x/tokenizer/bpe" - - "go.jknobloc.com/x/shelf" -) - -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/knobloch/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) { - d := dict.NewDict[*entropy.Entry]() - - // TODO we could count non empty lines to pre alloc dict - - scanner := bufio.NewScanner(r) - - for scanner.Scan() { - line := scanner.Text() - - if err := scanner.Err(); err != nil { - return nil, err - } - - var s string - var n int - - if _, err := fmt.Sscanf(line, "%s %d", &s, &n); err != nil { - return nil, err - } - - // TODO see cmd/main.go we currently need mbpe tok ?? - - if n < 100 { - continue - } - - ids := tok.Tokenize(s) - - if len(ids) < 4 { - continue - } - - b := to[int, uint8](ids) // TODO that conversion is fucking ass - - // TODO vs decode each buffer suffix and use s as map key - - d.Set(string(b), &entropy.Entry{ - Encoded: s, // TODO encoded or decoded - TokenIDs: b, - LogProbs: make([]float32, len(ids)), - N: 0, - }) - } - - return d, nil -} - -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/cmd/main.go b/research/entropy/cmd/main.go deleted file mode 100644 index cba8c33..0000000 --- a/research/entropy/cmd/main.go +++ /dev/null @@ -1,221 +0,0 @@ -package main - -import ( - "database/sql" - "fmt" - "log" - - _ "github.com/duckdb/duckdb-go/v2" - "go.jknobloc.com/x/research/entropy" - - "github.com/jonasknobloch/mbpe" - "go.jknobloc.com/x/dataset" - "go.jknobloc.com/x/gpt2" - "go.jknobloc.com/x/shelf" - "go.jknobloc.com/x/tokenizer/bpe" -) - -type NoPreTok struct{} - -func (p *NoPreTok) PreTokenize(s string) []string { - return []string{s} -} - -const CorpusMean = -0.7749100363022372 - -const Threshold = -2.0 - -func main() { - dict := mbpe.NewDict() - - if err := dict.Load(shelf.Abs("results/knobloch/minipile/dict.txt")); err != nil { - log.Fatal(err) - } - - if err := gpt2.InitializeEnvironment(); err != nil { - log.Fatal(err) - } - - // m := model() - // t := tokenizer() - // d := data() - // - // if err := entropy.Context(m, t, d); err != nil { - // log.Fatal(err) - // } - - 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) - - foo := tokenizerMBPE() - - type candidate struct { - src string - ids []int - means []float64 - occurrences int64 - } - - candidates := make([]candidate, 0) - baseline := entropy.NewBaseline() - - for i, c := range dict.Items()[10000:] { - if i > 10 { - break - } - - s := c.Src() - - if len(s) < 10 { - continue - } - - ids := foo.Tokenize(s) - - means, occurrences, err := entropy.MeanLogProbs(ids, db) - - if err != nil { - log.Fatal(err) - } - - if occurrences == 0 { - continue - } - - { - occurrences, err := entropy.LogProbs(ids, db) - - if err != nil { - log.Fatal(err) - } - - fmt.Println(ids) - fmt.Println(occurrences) - } - - candidates = append(candidates, candidate{s, ids, means, occurrences}) - - baseline.Add(means) - } - - for _, c := range candidates { - fmt.Println(c.ids) - fmt.Println(foo.Decoder().Decode([]string{c.src})) - fmt.Println(c.means, c.occurrences) - - fmt.Println("excess ", entropy.Excess(c.means, Threshold)) - fmt.Println("spikes ", entropy.Spikes(c.means)) - fmt.Println("deviations", baseline.Deviations(c.means)) - - fmt.Println() - fmt.Println() - fmt.Println() - } - - if err := gpt2.DestroyEnvironment(); err != nil { - log.Fatal(err) - } -} - -// func must[T any](v T, err error) T { -// if err != nil { -// log.Fatal(err) -// } -// -// return v -// } - -func model() *gpt2.Model { - cfg := gpt2.ConfigDefault() - - cfg.VocabSize = 256 - - opts := gpt2.Options{ - WithCache: false, - WithLogits: false, - WithLogProbs: true, - } - - m := gpt2.NewModel(shelf.Abs("models/mbpe/minipile-byte/gpt2_256_m000_minipile_seed/model_eval.onnx"), cfg, opts) - - if err := m.Init(); err != nil { - log.Fatal(err) - } - - return m -} - -func tokenizer() *bpe.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 - } - - return tok -} - -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 -} - -func data() dataset.Reader { - var miniPile *dataset.ParquetReader - - if r, err := dataset.NewParquetReader(shelf.Abs("data/minipile/validation")); err != nil { - log.Fatal(err) - } else { - miniPile = r - } - - return dataset.NewClampedReader(miniPile, 10) -} - -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 -} 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/entry.go b/research/entropy/entry.go index 746e34f..3ad2a26 100644 --- a/research/entropy/entry.go +++ b/research/entropy/entry.go @@ -4,5 +4,6 @@ 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 index 1e4814c..b5efad9 100644 --- a/research/entropy/experiment.go +++ b/research/entropy/experiment.go @@ -3,10 +3,22 @@ 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 { @@ -20,28 +32,36 @@ func Run(db *sql.DB, d *dict.Dict[*Entry]) error { // // TODO handle // } - rows, err := db.Query(`SELECT * FROM context`) + var rows *sql.Rows - if err != nil { + if r, err := db.Query(`SELECT * FROM context`); err != nil { return err + } else { + rows = r } defer rows.Close() - buffer := make([]uint8, 0, BufferSize) + bufferTokens := make([]uint8, 0, BufferSize) bufferLogProbs := make([]float32, 0, BufferSize) doc := -1 - maxTokens := 10000000 - currentToken := 0 + 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 currentToken == maxTokens { - // break + if currentTokens == maxTokens { + break } - currentToken++ + currentTokens++ var uid int var token int @@ -52,22 +72,22 @@ func Run(db *sql.DB, d *dict.Dict[*Entry]) error { return err } - if uid != doc || len(buffer) == BufferSize { - buffer = buffer[:0] + if uid != doc || len(bufferTokens) == BufferSize { + bufferTokens = bufferTokens[:0] bufferLogProbs = bufferLogProbs[:0] doc = uid } - buffer = append(buffer, uint8(token)) + bufferTokens = append(bufferTokens, uint8(token)) bufferLogProbs = append(bufferLogProbs, logProb) - n := len(buffer) + n := len(bufferTokens) // fmt.Print(n) for i := range n { - entry, ok := d.GetBytes(buffer[i:n]) + entry, ok := d.GetBytes(bufferTokens[i:n]) if !ok { continue @@ -79,6 +99,8 @@ func Run(db *sql.DB, d *dict.Dict[*Entry]) error { entry.N++ } + + pb.Add(1) } if err := rows.Err(); err != nil { @@ -90,30 +112,99 @@ func Run(db *sql.DB, d *dict.Dict[*Entry]) error { return cmp.Compare(b.Val.N, a.Val.N) }) - // for _, e := range d.Values() { - // if e.N < 100 { - // break - // } - // - // fmt.Println(e.Encoded) - // fmt.Println(spikes(e.LogProbs)) - // } + pb.Close() + + for _, e := range d.Values() { + spikes(e) + } + + if err := serializeSegments(d); err != nil { + return err + } return nil } -func spikes(logProbs []float32) []bool { - b := make([]bool, len(logProbs)) +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 i := range len(logProbs) { + 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 } - if d := logProbs[i-1] - logProbs[i]; d > 0 { - b[i] = true // TODO threshold + r[e.Encoded] = segmentsEntry{ + Value: segments, + OK: true, } } - return b + 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/segment.go b/research/entropy/segment.go deleted file mode 100644 index bdab576..0000000 --- a/research/entropy/segment.go +++ /dev/null @@ -1,346 +0,0 @@ -package entropy - -import ( - "database/sql" - "math" -) - -// CREATE TABLE context(uid INTEGER, token INTEGER, logprob FLOAT, pos INTEGER) -// -// sqlMeanLogProbs matches an exact token id sequence of any length, bound as a -// single list parameter, and returns one row per position: the position, the -// mean logprob over all occurrences, and the occurrence count. -// -// seq unrolls the parameter into (position, token) pairs. Joining context on -// token alone gives every position that matches some element of the sequence; -// pos - i + 1 is the start a match would have to begin at for that element to -// line up, so counting rows per (uid, start) says how many elements agree -// there. A start where all of them agree is an occurrence. -const sqlMeanLogProbs = ` -WITH seq AS ( - SELECT i, s[i] AS token - FROM (SELECT ?::INTEGER[] AS s) t, range(1, len(s) + 1) r(i) -), -cand AS ( - SELECT s.i AS i, - c.logprob AS logprob, - count(*) OVER (PARTITION BY c.uid, c.pos - s.i + 1) AS matched - FROM context c - JOIN seq s ON c.token = s.token -) -SELECT i, avg(logprob) AS mean, count(*) AS occurrences -FROM cand -WHERE matched = (SELECT count(*) FROM seq) -GROUP BY i -ORDER BY i` - -// sqlLogProbs matches the same sequences as sqlMeanLogProbs but skips the -// aggregation, returning one row per position per occurrence. -const sqlLogProbs = ` -WITH seq AS ( - SELECT i, s[i] AS token - FROM (SELECT ?::INTEGER[] AS s) t, range(1, len(s) + 1) r(i) -), -cand AS ( - SELECT c.uid AS uid, - c.pos - s.i + 1 AS start, - s.i AS i, - c.logprob AS logprob, - count(*) OVER (PARTITION BY c.uid, c.pos - s.i + 1) AS matched - FROM context c - JOIN seq s ON c.token = s.token -) -SELECT uid, start, i, logprob -FROM cand -WHERE matched = (SELECT count(*) FROM seq) -ORDER BY uid, start, i` - -// Occurrence is one match of a sequence, holding the logprob of every position -// in the context it was found in. -type Occurrence struct { - Document int - Offset int - Values []float64 -} - -// LogProbs returns every occurrence of the exact id sequence separately, -// leaving the contexts uncombined. -func LogProbs(ids []int, db *sql.DB) ([]Occurrence, error) { - seq := make([]int32, len(ids)) // the driver panics on []int - - for i, id := range ids { - seq[i] = int32(id) - } - - rows, err := db.Query(sqlLogProbs, seq) - - if err != nil { - return nil, err - } - - defer rows.Close() - - r := make([]Occurrence, 0) - - for rows.Next() { - var document, offset, pos int - var value float64 - - if err := rows.Scan(&document, &offset, &pos, &value); err != nil { - return nil, err - } - - if pos == 1 { - r = append(r, Occurrence{ - Document: document, - Offset: offset, - Values: make([]float64, 0, len(ids)), - }) - } - - if len(r) == 0 { - continue - } - - last := &r[len(r)-1] - - last.Values = append(last.Values, value) - } - - if err := rows.Err(); err != nil { - return nil, err - } - - return r, nil -} - -// MeanLogProbs returns the mean logprob per position over all occurrences of -// the exact id sequence, together with the number of occurrences. Occurrences -// of zero yield a nil slice: the sequence was never seen. -func MeanLogProbs(ids []int, db *sql.DB) ([]float64, int64, error) { - seq := make([]int32, len(ids)) // the driver panics on []int - - for i, id := range ids { - seq[i] = int32(id) - } - - rows, err := db.Query(sqlMeanLogProbs, seq) - - if err != nil { - return nil, 0, err - } - - defer rows.Close() - - values := make([]float64, 0, len(ids)) - - var occurrences int64 - - for rows.Next() { - var pos int - var mean float64 - var n int64 - - if err := rows.Scan(&pos, &mean, &n); err != nil { - return nil, 0, err - } - - values = append(values, mean) - - occurrences = n - } - - if err := rows.Err(); err != nil { - return nil, 0, err - } - - if len(values) == 0 { - return nil, 0, nil - } - - return values, occurrences, nil -} - -// Excess weights each position by how far it falls below threshold, in nats. -// Fires wherever surprisal is high in absolute terms, which in practice means -// the uncertain zone at the start of a chunk. -func Excess(values []float64, threshold float64) []float64 { - r := make([]float64, len(values)) - - for i, v := range values { - if d := threshold - v; d > 0 { - r[i] = d - } - } - - return r -} - -// Spikes weights each position by how much more surprising it is than the one -// before it, in nats. Catches boundaries an absolute threshold misses, since -// surprisal decays inside a unit and a boundary shows up as uncertainty rising -// again, however low the preceding tail got. Position 0 has no predecessor and -// is always 0. -func Spikes(values []float64) []float64 { - r := make([]float64, len(values)) - - for i := 1; i < len(values); i++ { - if d := values[i-1] - values[i]; d > 0 { - r[i] = d - } - } - - return r -} - -// Baseline accumulates mean logprobs per position index across sequences. -// Surprisal decays sharply with position -- onsets are uncertain, tails are -// forced -- so the raw value at a position says little on its own. Deviations -// measures how far a sequence departs from what its positions usually look -// like, which also cancels the downward pull the exact-sequence lookup puts on -// interior positions. -type Baseline struct { - sum []float64 - sumSq []float64 - noise []float64 - counts []int -} - -// MinVarianceShare floors the noise correction: however much sampling noise is -// estimated, a position keeps at least this share of its observed variance. -// Without it a position where every candidate agrees would correct to near-zero -// variance and produce enormous deviations from nothing. -const MinVarianceShare = 0.1 - -func NewBaseline() *Baseline { - return &Baseline{} -} - -func (b *Baseline) grow(n int) { - for len(b.counts) < n { - b.sum = append(b.sum, 0) - b.sumSq = append(b.sumSq, 0) - b.noise = append(b.noise, 0) - b.counts = append(b.counts, 0) - } -} - -func (b *Baseline) Add(values []float64) { - b.grow(len(values)) - - for i, v := range values { - b.sum[i] += v - b.sumSq[i] += v * v - b.counts[i]++ - } -} - -// AddCandidate is Add for a candidate whose means were estimated from n -// occurrences with the given per-position variances across those occurrences. -// The spread of candidate means overstates how much candidates really differ, -// by the sampling variance of the means themselves; recording it lets stats -// subtract it instead of passing the inflation on to every deviation. -func (b *Baseline) AddCandidate(values, variances []float64, n int) { - b.Add(values) - - if n < 2 { - return - } - - for i, v := range variances { - if i >= len(b.noise) { - break - } - - b.noise[i] += v / float64(n) - } -} - -func (b *Baseline) stats(pos int) (float64, float64, bool) { - if pos >= len(b.counts) || b.counts[pos] < 2 { - return 0, 0, false - } - - n := float64(b.counts[pos]) - mean := b.sum[pos] / n - - variance := (b.sumSq[pos] - n*mean*mean) / (n - 1) - - if variance <= 0 { - return mean, 0, false - } - - if corrected := variance - b.noise[pos]/n; corrected > variance*MinVarianceShare { - variance = corrected - } else { - variance = variance * MinVarianceShare - } - - return mean, math.Sqrt(variance), true -} - -// Summarize reduces the occurrences of one candidate to per-position means and -// variances, the inputs AddCandidate needs. -func Summarize(occurrences []Occurrence) ([]float64, []float64, int) { - if len(occurrences) == 0 { - return nil, nil, 0 - } - - n := len(occurrences[0].Values) - - means := make([]float64, n) - variances := make([]float64, n) - - for _, o := range occurrences { - for i, v := range o.Values { - if i >= n { - break - } - - means[i] += v - variances[i] += v * v - } - } - - for i := range means { - sum := means[i] - sumSq := variances[i] - - means[i] = sum / float64(len(occurrences)) - - if len(occurrences) < 2 { - variances[i] = 0 - - continue - } - - variances[i] = (sumSq - float64(len(occurrences))*means[i]*means[i]) / float64(len(occurrences)-1) - - if variances[i] < 0 { - variances[i] = 0 - } - } - - return means, variances, len(occurrences) -} - -// Deviations scores each position by how many standard deviations more -// surprising it is than the same position usually is. Positions at or below -// their positional mean score 0, as they carry no evidence of a boundary. -func (b *Baseline) Deviations(values []float64) []float64 { - r := make([]float64, len(values)) - - for i, v := range values { - mean, std, ok := b.stats(i) - - if !ok { - continue - } - - if d := (mean - v) / std; d > 0 { - r[i] = d - } - } - - return r -} diff --git a/research/entropy/segmenter.go b/research/entropy/segmenter.go deleted file mode 100644 index 1a15111..0000000 --- a/research/entropy/segmenter.go +++ /dev/null @@ -1,94 +0,0 @@ -package entropy - -import ( - "database/sql" -) - -type Tokenizer interface { - Tokenize(string) []int -} - -// Criterion turns per-position mean logprobs into per-position boundary -// weights, zero wherever there is no evidence of a boundary. -type Criterion func([]float64) []float64 - -// ExcessCriterion is Excess bound to a threshold. Weights come out in nats -// below that threshold, so a Segmenter using it wants a cutoff of 0. -func ExcessCriterion(threshold float64) Criterion { - return func(values []float64) []float64 { - return Excess(values, threshold) - } -} - -// Segmenter turns surprisal into segments, satisfying the mbpe Segmenter -// interface so it can stand in for Morfessor without touching the trainer. -// -// A boundary is placed wherever the criterion weighs a position above cutoff. -// Note that cutoff applies to the weights, not to the logprobs: with Spikes it -// is a minimum rise in nats and belongs above zero, with ExcessCriterion the -// threshold is already baked in and cutoff belongs at zero. -type Segmenter struct { - tokenizer Tokenizer - db *sql.DB - criterion Criterion - cutoff float64 -} - -func NewSegmenter(tokenizer Tokenizer, db *sql.DB, criterion Criterion, cutoff float64) *Segmenter { - return &Segmenter{ - tokenizer: tokenizer, - db: db, - criterion: criterion, - cutoff: cutoff, - } -} - -// Segment reports the segmentation of compound and whether surprisal was known -// for it at all. Compounds the model never saw come back unsegmented and not -// ok, which the trainer already treats as a reason to drop alpha to zero. -// -// The leading whitespace marker the trainer strips is put back before the -// lookup: the logprobs were measured on running text, so the first character -// of a word is only predictable given the space in front of it. -func (s *Segmenter) Segment(compound string) ([]string, bool) { - runes := []rune(compound) - - if len(runes) < 2 { - return []string{compound}, false - } - - ids := s.tokenizer.Tokenize("Ġ" + compound) - - if len(ids) != len(runes)+1 { - return []string{compound}, false // not one token per rune, cannot align - } - - values, occurrences, err := MeanLogProbs(ids, s.db) - - if err != nil || occurrences == 0 || len(values) != len(ids) { - return []string{compound}, false - } - - weights := s.criterion(values) - - segments := make([]string, 0, 4) - - start := 0 - - // ids[0] is the whitespace marker and ids[1] the first rune of compound, so - // a boundary there is the start of the chunk and carries no information. - // Skipping it also drops the onset spike, which dominates every sequence. - for i := 2; i < len(ids); i++ { - if weights[i] <= s.cutoff { - continue - } - - segments = append(segments, string(runes[start:i-1])) - - start = i - 1 - } - - segments = append(segments, string(runes[start:])) - - return segments, true -} -- cgit v1.2.3