diff options
| -rw-r--r-- | go.work | 1 | ||||
| -rw-r--r-- | research/entropy/cmd/dict3/main.go | 159 | ||||
| -rw-r--r-- | research/entropy/cmd/main.go | 221 | ||||
| -rw-r--r-- | research/entropy/context.go | 53 | ||||
| -rw-r--r-- | research/entropy/entry.go | 8 | ||||
| -rw-r--r-- | research/entropy/experiment.go | 119 | ||||
| -rw-r--r-- | research/entropy/go.mod | 32 | ||||
| -rw-r--r-- | research/entropy/go.sum | 36 | ||||
| -rw-r--r-- | research/entropy/segment.go | 346 | ||||
| -rw-r--r-- | research/entropy/segmenter.go | 94 | ||||
| -rw-r--r-- | tokenizer/bpe/utility.go | 4 |
11 files changed, 1073 insertions, 0 deletions
@@ -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) |
