summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--go.work1
-rw-r--r--research/entropy/cmd/dict3/main.go159
-rw-r--r--research/entropy/cmd/main.go221
-rw-r--r--research/entropy/context.go53
-rw-r--r--research/entropy/entry.go8
-rw-r--r--research/entropy/experiment.go119
-rw-r--r--research/entropy/go.mod32
-rw-r--r--research/entropy/go.sum36
-rw-r--r--research/entropy/segment.go346
-rw-r--r--research/entropy/segmenter.go94
-rw-r--r--tokenizer/bpe/utility.go4
11 files changed, 1073 insertions, 0 deletions
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)