summaryrefslogtreecommitdiff
path: root/research/entropy
diff options
context:
space:
mode:
Diffstat (limited to 'research/entropy')
-rw-r--r--research/entropy/cmd/segment/main.go134
-rw-r--r--research/entropy/context.go53
-rw-r--r--research/entropy/entry.go9
-rw-r--r--research/entropy/experiment.go210
-rw-r--r--research/entropy/go.mod32
-rw-r--r--research/entropy/go.sum36
6 files changed, 474 insertions, 0 deletions
diff --git a/research/entropy/cmd/segment/main.go b/research/entropy/cmd/segment/main.go
new file mode 100644
index 0000000..cc6aea4
--- /dev/null
+++ b/research/entropy/cmd/segment/main.go
@@ -0,0 +1,134 @@
+package main
+
+import (
+ "database/sql"
+ "fmt"
+ "io"
+ "log"
+ "os"
+
+ "github.com/jonasknobloch/mbpe"
+
+ "go.jknobloc.com/x/dict"
+ "go.jknobloc.com/x/research/entropy"
+ "go.jknobloc.com/x/shelf"
+ "go.jknobloc.com/x/tokenizer/bpe"
+
+ _ "github.com/duckdb/duckdb-go/v2"
+)
+
+type NoPreTok struct{}
+
+func (p *NoPreTok) PreTokenize(s string) []string {
+ return []string{s}
+}
+
+func main() {
+ o := shelf.Abs("results/runs-entropy/lesci_minipile_test_256_seed_default/gpt2_256_m000_minipile_seed/lesci.db")
+
+ var db *sql.DB
+
+ if database, err := initDatabase(o); err != nil {
+ log.Fatal(err)
+ } else {
+ db = database
+ }
+
+ defer db.Close()
+
+ db.SetMaxOpenConns(1)
+
+ var file *os.File
+
+ if f, err := os.Open(shelf.Abs("results/entropy/minipile/dict.txt")); err != nil {
+ log.Fatal(err)
+ } else {
+ file = f
+ }
+
+ tok := tokenizerMBPE()
+
+ var dict *dict.Dict[*entropy.Entry]
+
+ if d, err := loadDict(file, tok); err != nil {
+ log.Fatal(err)
+ } else {
+ dict = d
+ }
+
+ if err := entropy.Run(db, dict); err != nil {
+ log.Fatal(err)
+ }
+}
+
+func initDatabase(dsn string) (*sql.DB, error) {
+ var db *sql.DB
+
+ if database, err := sql.Open("duckdb", dsn); err != nil {
+ return nil, err
+ } else {
+ db = database
+ }
+
+ if err := db.Ping(); err != nil {
+ _ = db.Close()
+
+ return nil, err
+ }
+
+ return db, nil
+}
+
+func loadDict(r io.Reader, tok *mbpe.Tokenizer) (*dict.Dict[*entropy.Entry], error) {
+ return dict.LoadDict[*entropy.Entry](r, func(l string) (string, *entropy.Entry) {
+ var s string
+ var n int
+
+ if _, err := fmt.Sscanf(l, "%s %d", &s, &n); err != nil {
+ panic(err)
+ }
+
+ ids := tok.Tokenize(s)
+
+ b := to[int, uint8](ids) // TODO that conversion is fucking ass
+
+ return string(b), &entropy.Entry{
+ Encoded: s, // TODO encoded or decoded
+ TokenIDs: b,
+ LogProbs: make([]float32, len(ids)),
+ Bounds: make([]bool, len(ids)+1),
+ N: 0,
+ }
+ })
+}
+
+func to[a bpe.Integer, b bpe.Integer](s []a) []b {
+ r := make([]b, len(s))
+
+ // TODO handle overflow
+
+ for i, v := range s {
+ r[i] = b(v)
+ }
+
+ return r
+}
+
+func tokenizerMBPE() *mbpe.Tokenizer {
+ var tok *bpe.Tokenizer
+
+ v := shelf.Abs("tokenizers/minipile/tokenizer_gpt2_256_m000_minipile/vocab.json")
+ m := shelf.Abs("tokenizers/minipile/tokenizer_gpt2_256_m000_minipile/merges.txt")
+
+ if t, err := bpe.NewTokenizerFromFiles(v, m, bpe.DefaultConfig()); err != nil {
+ log.Fatal(err)
+ } else {
+ tok = t
+ }
+
+ foo := bpe.MBPE(tok)
+
+ foo.SetPreTokenizer(&NoPreTok{})
+
+ return foo
+}
diff --git a/research/entropy/context.go b/research/entropy/context.go
new file mode 100644
index 0000000..3993c03
--- /dev/null
+++ b/research/entropy/context.go
@@ -0,0 +1,53 @@
+package entropy
+
+import (
+ "fmt"
+
+ "go.jknobloc.com/x/dataset"
+ "go.jknobloc.com/x/llm"
+)
+
+type logProb struct {
+ document int
+ token int
+ value float32
+ offset int
+}
+
+func Context(model llm.Causal, tokenizer llm.Tokenizer, data dataset.Reader) error {
+ evaluatorConfig := llm.EvaluatorConfig{
+ BatchSize: 32,
+ NumWorkers: 16,
+ }
+
+ tokenBufferConfig := llm.TokenBufferConfig{
+ Window: 1024,
+ Stride: 512,
+ PadLeft: false,
+ PadRight: false,
+ PadTokenID: 256,
+ }
+
+ eval := llm.NewEvaluator(model, tokenizer, func(job llm.Job, logProbs []float32, tokens []int) []logProb {
+ r := make([]logProb, len(tokens))
+
+ for i, token := range tokens {
+ r[i] = logProb{
+ document: job.Document,
+ token: token,
+ value: logProbs[i],
+ offset: job.Position*tokenBufferConfig.Stride + job.Seen + i,
+ }
+ }
+
+ return r
+ }, evaluatorConfig)
+
+ return eval.RunAndCollect("Context", data, tokenBufferConfig, func(r []logProb) error {
+ for _, l := range r {
+ fmt.Println(l) // TODO implement
+ }
+
+ return nil
+ })
+}
diff --git a/research/entropy/entry.go b/research/entropy/entry.go
new file mode 100644
index 0000000..3ad2a26
--- /dev/null
+++ b/research/entropy/entry.go
@@ -0,0 +1,9 @@
+package entropy
+
+type Entry struct {
+ Encoded string
+ TokenIDs []uint8
+ LogProbs []float32
+ Bounds []bool
+ N int
+}
diff --git a/research/entropy/experiment.go b/research/entropy/experiment.go
new file mode 100644
index 0000000..b5efad9
--- /dev/null
+++ b/research/entropy/experiment.go
@@ -0,0 +1,210 @@
+package entropy
+
+import (
+ "cmp"
+ "database/sql"
+ "encoding/gob"
+ "os"
+ "time"
+
+ "github.com/jonasknobloch/mbpe"
+
+ "go.jknobloc.com/x/dict"
+ "go.jknobloc.com/x/shelf"
+ "go.jknobloc.com/x/tui"
+)
+
+type segmentsEntry struct {
+ Value []string
+ OK bool
+}
+
+const BufferSize = 1024
+
+func Run(db *sql.DB, d *dict.Dict[*Entry]) error {
+ // m := 0
+ //
+ // for _, v := range d.Values() {
+ // m = max(m, len(v.TokenIDs))
+ // }
+ //
+ // if m > BufferSize {
+ // // TODO handle
+ // }
+
+ var rows *sql.Rows
+
+ if r, err := db.Query(`SELECT * FROM context`); err != nil {
+ return err
+ } else {
+ rows = r
+ }
+
+ defer rows.Close()
+
+ bufferTokens := make([]uint8, 0, BufferSize)
+ bufferLogProbs := make([]float32, 0, BufferSize)
+
+ doc := -1
+
+ maxTokens := 100000
+ currentTokens := 0
+
+ pb := tui.NewProgressBar("Adding LogProbs", 20, 58508191, time.Now())
+
+ pb.Start(1*time.Second, func() int {
+ return pb.Completed()
+ })
+
+ for rows.Next() {
+ if currentTokens == maxTokens {
+ break
+ }
+
+ currentTokens++
+
+ var uid int
+ var token int
+ var logProb float32
+ var pos int
+
+ if err := rows.Scan(&uid, &token, &logProb, &pos); err != nil {
+ return err
+ }
+
+ if uid != doc || len(bufferTokens) == BufferSize {
+ bufferTokens = bufferTokens[:0]
+ bufferLogProbs = bufferLogProbs[:0]
+
+ doc = uid
+ }
+
+ bufferTokens = append(bufferTokens, uint8(token))
+ bufferLogProbs = append(bufferLogProbs, logProb)
+
+ n := len(bufferTokens)
+
+ // fmt.Print(n)
+
+ for i := range n {
+ entry, ok := d.GetBytes(bufferTokens[i:n])
+
+ if !ok {
+ continue
+ }
+
+ for j := range entry.LogProbs {
+ entry.LogProbs[j] += bufferLogProbs[j]
+ }
+
+ entry.N++
+ }
+
+ pb.Add(1)
+ }
+
+ if err := rows.Err(); err != nil {
+ return err
+ }
+
+ // debug
+ d.Sort(func(a, b dict.Entry[*Entry]) int {
+ return cmp.Compare(b.Val.N, a.Val.N)
+ })
+
+ pb.Close()
+
+ for _, e := range d.Values() {
+ spikes(e)
+ }
+
+ if err := serializeSegments(d); err != nil {
+ return err
+ }
+
+ return nil
+}
+
+func spikes(e *Entry) {
+ for i := range len(e.LogProbs) {
+ if i == 0 {
+ continue
+ }
+
+ if i == 1 && e.TokenIDs[0] == 220 {
+ continue
+ }
+
+ // TODO threshold
+
+ if d := e.LogProbs[i-1] - e.LogProbs[i]; d > 0 {
+ e.Bounds[i] = true
+ }
+ }
+}
+
+func serializeSegments(d *dict.Dict[*Entry]) error {
+ r := make(map[string]segmentsEntry, d.Len())
+
+ for _, e := range d.Values() {
+ segments := make([]string, 0)
+
+ decoded := decode(e.Encoded)
+
+ i := 0
+
+ for j, b := range e.Bounds {
+ if b {
+ segments = append(segments, encode(decoded[i:j]))
+
+ i = j
+ }
+ }
+
+ segments = append(segments, encode(decoded[i:])) // TODO or just use OK false?
+
+ // debug
+ if i == 0 {
+ continue
+ }
+
+ r[e.Encoded] = segmentsEntry{
+ Value: segments,
+ OK: true,
+ }
+ }
+
+ var file *os.File
+
+ if f, err := os.Create(shelf.Abs("results/entropy/minipile/segments.gob")); err != nil {
+ return err
+ } else {
+ file = f
+ }
+
+ defer file.Close()
+
+ enc := gob.NewEncoder(file)
+
+ return enc.Encode(r)
+}
+
+func encode(s string) string {
+ result := ""
+
+ for _, b := range []byte(s) {
+ result += mbpe.BytesChar[b]
+ }
+
+ return result
+}
+
+func decode(s string) string {
+ result := make([]byte, 0)
+
+ for _, r := range s {
+ result = append(result, mbpe.CharBytes[string(r)])
+ }
+
+ return string(result)
+}
diff --git a/research/entropy/go.mod b/research/entropy/go.mod
new file mode 100644
index 0000000..ab7ba3d
--- /dev/null
+++ b/research/entropy/go.mod
@@ -0,0 +1,32 @@
+module go.jknobloc.com/x/research/entropy
+
+go 1.25.0
+
+require github.com/duckdb/duckdb-go/v2 v2.10501.0
+
+require (
+ github.com/apache/arrow-go/v18 v18.5.1 // indirect
+ github.com/duckdb/duckdb-go-bindings v0.10501.0 // indirect
+ github.com/duckdb/duckdb-go-bindings/lib/darwin-amd64 v0.10501.0 // indirect
+ github.com/duckdb/duckdb-go-bindings/lib/darwin-arm64 v0.10501.0 // indirect
+ github.com/duckdb/duckdb-go-bindings/lib/linux-amd64 v0.10501.0 // indirect
+ github.com/duckdb/duckdb-go-bindings/lib/linux-arm64 v0.10501.0 // indirect
+ github.com/duckdb/duckdb-go-bindings/lib/windows-amd64 v0.10501.0 // indirect
+ github.com/go-viper/mapstructure/v2 v2.5.0 // indirect
+ github.com/goccy/go-json v0.10.5 // indirect
+ github.com/google/flatbuffers v25.12.19+incompatible // indirect
+ github.com/google/go-cmp v0.7.0 // indirect
+ github.com/google/uuid v1.6.0 // indirect
+ github.com/klauspost/compress v1.18.3 // indirect
+ github.com/klauspost/cpuid/v2 v2.3.0 // indirect
+ github.com/pierrec/lz4/v4 v4.1.25 // indirect
+ github.com/zeebo/xxh3 v1.1.0 // indirect
+ golang.org/x/exp v0.0.0-20260112195511-716be5621a96 // indirect
+ golang.org/x/mod v0.33.0 // indirect
+ golang.org/x/sync v0.20.0 // indirect
+ golang.org/x/sys v0.41.0 // indirect
+ golang.org/x/telemetry v0.0.0-20260209163413-e7419c687ee4 // indirect
+ golang.org/x/tools v0.42.0 // indirect
+ golang.org/x/xerrors v0.0.0-20240903120638-7835f813f4da // indirect
+ gonum.org/v1/gonum v0.17.0 // indirect
+)
diff --git a/research/entropy/go.sum b/research/entropy/go.sum
new file mode 100644
index 0000000..70ff7c1
--- /dev/null
+++ b/research/entropy/go.sum
@@ -0,0 +1,36 @@
+github.com/andybalholm/brotli v1.2.0 h1:ukwgCxwYrmACq68yiUqwIWnGY0cTPox/M94sVwToPjQ=
+github.com/apache/arrow-go/v18 v18.5.1 h1:yaQ6zxMGgf9YCYw4/oaeOU3AULySDlAYDOcnr4LdHdI=
+github.com/apache/thrift v0.22.0 h1:r7mTJdj51TMDe6RtcmNdQxgn9XcyfGDOzegMDRg47uc=
+github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM=
+github.com/duckdb/duckdb-go-bindings v0.10501.0 h1:BR21HkcALr9Lm+Ios2vEPaaB5oRRxGJHONzkS0bnOKE=
+github.com/duckdb/duckdb-go-bindings/lib/darwin-amd64 v0.10501.0 h1:InnDiz/iBHUzwI/4xkigTq6PRrIx+9L+eC2NfCShgWc=
+github.com/duckdb/duckdb-go-bindings/lib/darwin-arm64 v0.10501.0 h1:XLMUi/9QJcN8Bp77ML/QPwynX8f9RAg4VUiTdPzRUEU=
+github.com/duckdb/duckdb-go-bindings/lib/linux-amd64 v0.10501.0 h1:td84w8XucSPQoxGC84RYIxTu1+RV+fjeFIHVk3GLXog=
+github.com/duckdb/duckdb-go-bindings/lib/linux-arm64 v0.10501.0 h1:XLw03uWhdQvAFU6unJ2MtQvSFi4LNCfhiOM644MyAHw=
+github.com/duckdb/duckdb-go-bindings/lib/windows-amd64 v0.10501.0 h1:jhhOonew2VOcTN4f+BOlnOTUgHEp797Uee2Tq8xMno8=
+github.com/duckdb/duckdb-go/v2 v2.10501.0 h1:vYgvKBfotrZqpBESqHXYF5NVlbaYHRf0VrQEXtb/jnU=
+github.com/go-viper/mapstructure/v2 v2.5.0 h1:vM5IJoUAy3d7zRSVtIwQgBj7BiWtMPfmPEgAXnvj1Ro=
+github.com/goccy/go-json v0.10.5 h1:Fq85nIqj+gXn/S5ahsiTlK3TmC85qgirsdTP/+DeaC4=
+github.com/golang/snappy v1.0.0 h1:Oy607GVXHs7RtbggtPBnr2RmDArIsAefDwvrdWvRhGs=
+github.com/google/flatbuffers v25.12.19+incompatible h1:haMV2JRRJCe1998HeW/p0X9UaMTK6SDo0ffLn2+DbLs=
+github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
+github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
+github.com/klauspost/asmfmt v1.3.2 h1:4Ri7ox3EwapiOjCki+hw14RyKk201CN4rzyCJRFLpK4=
+github.com/klauspost/compress v1.18.3 h1:9PJRvfbmTabkOX8moIpXPbMMbYN60bWImDDU7L+/6zw=
+github.com/klauspost/cpuid/v2 v2.3.0 h1:S4CRMLnYUhGeDFDqkGriYKdfoFlDnMtqTiI/sFzhA9Y=
+github.com/minio/asm2plan9s v0.0.0-20200509001527-cdd76441f9d8 h1:AMFGa4R4MiIpspGNG7Z948v4n35fFGB3RR3G/ry4FWs=
+github.com/minio/c2goasm v0.0.0-20190812172519-36a3d3bbc4f3 h1:+n/aFZefKZp7spd8DFdX7uMikMLXX4oubIzJF4kv/wI=
+github.com/pierrec/lz4/v4 v4.1.25 h1:kocOqRffaIbU5djlIBr7Wh+cx82C0vtFb0fOurZHqD0=
+github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 h1:Jamvg5psRIccs7FGNTlIRMkT8wgtp5eCXdBlqhYGL6U=
+github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
+github.com/zeebo/assert v1.3.0 h1:g7C04CbJuIDKNPFHmsk4hwZDO5O+kntRxzaUoNXj+IQ=
+github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs=
+golang.org/x/exp v0.0.0-20260112195511-716be5621a96 h1:Z/6YuSHTLOHfNFdb8zVZomZr7cqNgTJvA8+Qz75D8gU=
+golang.org/x/mod v0.33.0 h1:tHFzIWbBifEmbwtGz65eaWyGiGZatSrT9prnU8DbVL8=
+golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
+golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k=
+golang.org/x/telemetry v0.0.0-20260209163413-e7419c687ee4 h1:bTLqdHv7xrGlFbvf5/TXNxy/iUwwdkjhqQTJDjW7aj0=
+golang.org/x/tools v0.42.0 h1:uNgphsn75Tdz5Ji2q36v/nsFSfR/9BRFvqhGBaJGd5k=
+golang.org/x/xerrors v0.0.0-20240903120638-7835f813f4da h1:noIWHXmPHxILtqtCOPIhSt0ABwskkZKjD3bXGnZGpNY=
+gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
+gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=