diff options
Diffstat (limited to 'research/entropy')
| -rw-r--r-- | research/entropy/cmd/segment/main.go | 134 | ||||
| -rw-r--r-- | research/entropy/context.go | 53 | ||||
| -rw-r--r-- | research/entropy/entry.go | 9 | ||||
| -rw-r--r-- | research/entropy/experiment.go | 210 | ||||
| -rw-r--r-- | research/entropy/go.mod | 32 | ||||
| -rw-r--r-- | research/entropy/go.sum | 36 |
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= |
