diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-17 16:07:14 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-17 16:07:14 +0200 |
| commit | 3fe37224ccc2500e7e767e5fc00163ef61182828 (patch) | |
| tree | 7bd14e080c080b24026660beda42817a3c9dc77f /research/sander | |
| parent | 4d8f647bfb54899e3052931790e288da7af97050 (diff) | |
Add sander module
Diffstat (limited to 'research/sander')
| -rw-r--r-- | research/sander/analyze.go | 49 | ||||
| -rw-r--r-- | research/sander/cmd/sander/main.go | 42 | ||||
| -rw-r--r-- | research/sander/experiment.go | 85 | ||||
| -rw-r--r-- | research/sander/extract.go | 52 | ||||
| -rw-r--r-- | research/sander/go.mod | 61 | ||||
| -rw-r--r-- | research/sander/go.sum | 175 | ||||
| -rw-r--r-- | research/sander/plot.go | 100 | ||||
| -rw-r--r-- | research/sander/reference.go | 59 |
8 files changed, 623 insertions, 0 deletions
diff --git a/research/sander/analyze.go b/research/sander/analyze.go new file mode 100644 index 0000000..c58195d --- /dev/null +++ b/research/sander/analyze.go @@ -0,0 +1,49 @@ +package sander + +import ( + "context" + "database/sql" + "fmt" +) + +var ( + sqlSimilarity = ` + CREATE TABLE similarity_buffer AS + WITH centroid AS ( + SELECT list(avg_val ORDER BY idx)::FLOAT[768] AS vec + FROM ( + SELECT idx, avg(val) AS avg_val + FROM embeddings + CROSS JOIN LATERAL UNNEST(embedding::FLOAT[]) WITH ORDINALITY AS t(val, idx) + WHERE reference = true + GROUP BY idx + ) + ) + SELECT e.token_id, array_distance(e.embedding, c.vec) AS distance + FROM embeddings e, centroid c +` +) + +func (e *Experiment) Analyze(db *sql.DB) error { + ctx := context.Background() + + var tx *sql.Tx + + if t, err := db.BeginTx(ctx, nil); err != nil { + return err + } else { + tx = t + } + + defer tx.Rollback() + + if _, err := tx.ExecContext(ctx, `DROP TABLE IF EXISTS similarity_buffer`); err != nil { + return err + } + + if _, err := tx.ExecContext(ctx, sqlSimilarity); err != nil { + return fmt.Errorf("similarity: %w", err) + } + + return tx.Commit() +} diff --git a/research/sander/cmd/sander/main.go b/research/sander/cmd/sander/main.go new file mode 100644 index 0000000..fd9ece3 --- /dev/null +++ b/research/sander/cmd/sander/main.go @@ -0,0 +1,42 @@ +package main + +import ( + "fmt" + "log" + "path" + + "go.jknobloc.com/x/research/sander" + "go.jknobloc.com/x/tokenizer/bpe" +) + +func main() { + vocab := []string{"50256"} + alpha := []string{"m000", "m010", "m020", "m030", "m040", "m050", "m060", "m070", "m080", "m090", "m100"} + + for _, v := range vocab { + for _, a := range alpha { + src := fmt.Sprintf("gpt2/models/mbpe/gpt2_%s_%s_babylm_v2", v, a) + dst := fmt.Sprintf("out/sander/mbpe/gpt2_%s_%s_babylm_v2", v, a) + + if err := run(src, dst); err != nil { + log.Fatal(err) + } + } + } + +} + +func run(src, dst string) error { + t := must(bpe.NewTokenizerFromFiles(path.Join(src, "vocab.json"), path.Join(src, "merges.txt"))) + e := must(sander.NewExperiment(dst, path.Join(src, "model.onnx"), sander.UnusedTokensMBPE(t))) + + return e.Run() +} + +func must[T any](v T, err error) T { + if err != nil { + log.Fatal(err) + } + + return v +} diff --git a/research/sander/experiment.go b/research/sander/experiment.go new file mode 100644 index 0000000..1521cf8 --- /dev/null +++ b/research/sander/experiment.go @@ -0,0 +1,85 @@ +package sander + +import ( + "database/sql" + "fmt" + "log" + "os" + "path/filepath" + + _ "github.com/duckdb/duckdb-go/v2" +) + +type Experiment struct { + name string + model string + reference map[int64]struct{} +} + +func NewExperiment(name, model string, reference map[int64]struct{}) (*Experiment, error) { + e := &Experiment{ + name: name, + model: model, + reference: reference, + } + + return e, nil +} + +func (e *Experiment) Run() error { + if err := os.MkdirAll(e.name, 0775); err != nil { + log.Fatal(err) + } + + dsn := filepath.Join(e.name, "wte.db") + + var db *sql.DB + + if database, err := initDatabase(dsn); err != nil { + return err + } else { + db = database + } + + defer db.Close() + + db.SetMaxOpenConns(1) + + fmt.Println(e.name) + + if _, err := db.Exec(`INSTALL vss; LOAD vss;`); err != nil { + return err + } + + if err := e.Extract(db); err != nil { + return err + } + + if err := e.Analyze(db); err != nil { + return err + } + + if err := e.Plot(db); err != nil { + return err + } + + return nil +} + +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/sander/extract.go b/research/sander/extract.go new file mode 100644 index 0000000..de31b1f --- /dev/null +++ b/research/sander/extract.go @@ -0,0 +1,52 @@ +package sander + +import ( + "database/sql" + "database/sql/driver" + + "go.jknobloc.com/x/onnx" + "go.jknobloc.com/x/research/lesci" + "go.jknobloc.com/x/tensor" +) + +func (e *Experiment) Extract(db *sql.DB) error { + var data []float32 + var shape []int + + if d, s, err := onnx.ExtractInitializer(e.model, "transformer.wte.weight"); err != nil { + return err + } else { + data = d + shape = s + } + + if shape[1] != 768 { + panic("unimplemented") + } + + t := tensor.NewDense[float32](shape, data) + + if _, err := db.Exec(` + DROP TABLE IF EXISTS embeddings; + CREATE TABLE embeddings (token_id INTEGER, embedding FLOAT[768], reference BOOLEAN); + `); err != nil { + return err + } + + return lesci.AppendRows(db, "embeddings", func(appendFunc lesci.AppendFunc) error { + for i := range shape[0] { + embedding, ok := t.Select(0, i).Contiguous().Data() + + if !ok { + panic("") + } + + _, reference := e.reference[int64(i)] + + if err := appendFunc([]driver.Value{int32(i), embedding, reference}); err != nil { + return err + } + } + return nil + }) +} diff --git a/research/sander/go.mod b/research/sander/go.mod new file mode 100644 index 0000000..c1dc1eb --- /dev/null +++ b/research/sander/go.mod @@ -0,0 +1,61 @@ +module go.jknobloc.com/x/research/sander + +go 1.25.0 + +require ( + github.com/duckdb/duckdb-go/v2 v2.10501.0 + go.jknobloc.com/x/onnx v0.0.0-20260417134810-4d8f647bfb54 + go.jknobloc.com/x/research/lesci v0.0.0-20260414175339-402477b200d1 + go.jknobloc.com/x/tensor v0.0.0-20260414175339-402477b200d1 + go.jknobloc.com/x/tokenizer v0.0.0-20260410210408-456f8218a2d1 + gonum.org/v1/plot v0.16.0 +) + +require ( + codeberg.org/go-fonts/liberation v0.5.0 // indirect + codeberg.org/go-latex/latex v0.2.0 // indirect + codeberg.org/go-pdf/fpdf v0.11.1 // indirect + git.sr.ht/~sbinet/gg v0.7.0 // indirect + github.com/ajstarks/svgo v0.0.0-20211024235047-1546f124cd8b // indirect + github.com/andybalholm/brotli v1.2.0 // indirect + github.com/apache/arrow-go/v18 v18.5.1 // indirect + github.com/apache/thrift v0.22.0 // indirect + github.com/cespare/xxhash/v2 v2.3.0 // 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/golang/freetype v0.0.0-20170609003504-e2365dfdc4a0 // indirect + github.com/golang/snappy v1.0.0 // indirect + github.com/google/flatbuffers v25.12.19+incompatible // indirect + github.com/google/uuid v1.6.0 // indirect + github.com/jonasknobloch/mbpe v0.1.1 // indirect + github.com/klauspost/asmfmt v1.3.2 // indirect + github.com/klauspost/compress v1.18.3 // indirect + github.com/klauspost/cpuid/v2 v2.3.0 // indirect + github.com/minio/asm2plan9s v0.0.0-20200509001527-cdd76441f9d8 // indirect + github.com/minio/c2goasm v0.0.0-20190812172519-36a3d3bbc4f3 // indirect + github.com/pierrec/lz4/v4 v4.1.25 // indirect + github.com/zeebo/xxh3 v1.1.0 // indirect + go.jknobloc.com/x/dataset v0.0.0-20260410210408-456f8218a2d1 // indirect + go.jknobloc.com/x/llm v0.0.0-20260410210408-456f8218a2d1 // indirect + go.jknobloc.com/x/tui v0.0.0-20260324194423-87bbece7e040 // indirect + golang.org/x/exp v0.0.0-20260112195511-716be5621a96 // indirect + golang.org/x/image v0.37.0 // indirect + golang.org/x/mod v0.33.0 // indirect + golang.org/x/net v0.50.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/text v0.35.0 // 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 + google.golang.org/genproto/googleapis/rpc v0.0.0-20251029180050-ab9386a59fda // indirect + google.golang.org/grpc v1.78.0 // indirect + google.golang.org/protobuf v1.36.11 // indirect +) diff --git a/research/sander/go.sum b/research/sander/go.sum new file mode 100644 index 0000000..ba46cb2 --- /dev/null +++ b/research/sander/go.sum @@ -0,0 +1,175 @@ +codeberg.org/go-fonts/dejavu v0.4.0 h1:2yn58Vkh4CFK3ipacWUAIE3XVBGNa0y1bc95Bmfx91I= +codeberg.org/go-fonts/dejavu v0.4.0/go.mod h1:abni088lmhQJvso2Lsb7azCKzwkfcnttl6tL1UTWKzg= +codeberg.org/go-fonts/latin-modern v0.4.0 h1:vkRCc1y3whKA7iL9Ep0fSGVuJfqjix0ica9UflHORO8= +codeberg.org/go-fonts/latin-modern v0.4.0/go.mod h1:BF68mZznJ9QHn+hic9ks2DaFl4sR5YhfM6xTYaP9vNw= +codeberg.org/go-fonts/liberation v0.5.0 h1:SsKoMO1v1OZmzkG2DY+7ZkCL9U+rrWI09niOLfQ5Bo0= +codeberg.org/go-fonts/liberation v0.5.0/go.mod h1:zS/2e1354/mJ4pGzIIaEtm/59VFCFnYC7YV6YdGl5GU= +codeberg.org/go-latex/latex v0.2.0 h1:Ol/a6VHY06N+5gPfewswymoRb5ZcKDXWVaVegcx4hbI= +codeberg.org/go-latex/latex v0.2.0/go.mod h1:VJAwQir7/T8LZxj7xAPivISKiVOwkMpQ8bTuPQ31X0Y= +codeberg.org/go-pdf/fpdf v0.11.1 h1:U8+coOTDVLxHIXZgGvkfQEi/q0hYHYvEHFuGNX2GzGs= +codeberg.org/go-pdf/fpdf v0.11.1/go.mod h1:Y0DGRAdZ0OmnZPvjbMp/1bYxmIPxm0ws4tfoPOc4LjU= +git.sr.ht/~sbinet/cmpimg v0.1.0 h1:E0zPRk2muWuCqSKSVZIWsgtU9pjsw3eKHi8VmQeScxo= +git.sr.ht/~sbinet/cmpimg v0.1.0/go.mod h1:FU12psLbF4TfNXkKH2ZZQ29crIqoiqTZmeQ7dkp/pxE= +git.sr.ht/~sbinet/gg v0.7.0 h1:YmNf7YKd7diDMTPm86hZa1EM3pbkOyD/zzjl0LZUdNM= +git.sr.ht/~sbinet/gg v0.7.0/go.mod h1:VYeli15tpMM4EvqlivlVbbyvWZlOU+EZn4XZmfBGUdM= +github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= +github.com/ajstarks/deck v0.0.0-20200831202436-30c9fc6549a9/go.mod h1:JynElWSGnm/4RlzPXRlREEwqTHAN3T56Bv2ITsFT3gY= +github.com/ajstarks/deck/generate v0.0.0-20210309230005-c3f852c02e19/go.mod h1:T13YZdzov6OU0A1+RfKZiZN9ca6VeKdBdyDV+BY97Tk= +github.com/ajstarks/svgo v0.0.0-20211024235047-1546f124cd8b h1:slYM766cy2nI3BwyRiyQj/Ud48djTMtMebDqepE95rw= +github.com/ajstarks/svgo v0.0.0-20211024235047-1546f124cd8b/go.mod h1:1KcenG0jGWcpt8ov532z81sp/kMMUG485J2InIOyADM= +github.com/andybalholm/brotli v1.2.0 h1:ukwgCxwYrmACq68yiUqwIWnGY0cTPox/M94sVwToPjQ= +github.com/andybalholm/brotli v1.2.0/go.mod h1:rzTDkvFWvIrjDXZHkuS16NPggd91W3kUSvPlQ1pLaKY= +github.com/apache/arrow-go/v18 v18.5.1 h1:yaQ6zxMGgf9YCYw4/oaeOU3AULySDlAYDOcnr4LdHdI= +github.com/apache/arrow-go/v18 v18.5.1/go.mod h1:OCCJsmdq8AsRm8FkBSSmYTwL/s4zHW9CqxeBxEytkNE= +github.com/apache/thrift v0.22.0 h1:r7mTJdj51TMDe6RtcmNdQxgn9XcyfGDOzegMDRg47uc= +github.com/apache/thrift v0.22.0/go.mod h1:1e7J/O1Ae6ZQMTYdy9xa3w9k+XHWPfRvdPyJeynQ+/g= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= +github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM= +github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/dlclark/regexp2 v1.11.5 h1:Q/sSnsKerHeCkc/jSTNq1oCm7KiVgUMZRDUoRu0JQZQ= +github.com/dlclark/regexp2 v1.11.5/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8= +github.com/duckdb/duckdb-go-bindings v0.10501.0 h1:BR21HkcALr9Lm+Ios2vEPaaB5oRRxGJHONzkS0bnOKE= +github.com/duckdb/duckdb-go-bindings v0.10501.0/go.mod h1:UiTBFhbFLPI8+jX7hi3N577KlKOZGj/BW5qSN904658= +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-amd64 v0.10501.0/go.mod h1:EnAvZh1kNJHp5yF+M1ZHNEvapnmt6anq1xXHVrAGqMo= +github.com/duckdb/duckdb-go-bindings/lib/darwin-arm64 v0.10501.0 h1:XLMUi/9QJcN8Bp77ML/QPwynX8f9RAg4VUiTdPzRUEU= +github.com/duckdb/duckdb-go-bindings/lib/darwin-arm64 v0.10501.0/go.mod h1:IGLSeEcFhNeZF16aVjQCULD7TsFZKG5G7SyKJAXKp5c= +github.com/duckdb/duckdb-go-bindings/lib/linux-amd64 v0.10501.0 h1:td84w8XucSPQoxGC84RYIxTu1+RV+fjeFIHVk3GLXog= +github.com/duckdb/duckdb-go-bindings/lib/linux-amd64 v0.10501.0/go.mod h1:KAIynZ0GHCS7X5fRyuFnQMg/SZBPK/bS9OCOVojClxw= +github.com/duckdb/duckdb-go-bindings/lib/linux-arm64 v0.10501.0 h1:XLw03uWhdQvAFU6unJ2MtQvSFi4LNCfhiOM644MyAHw= +github.com/duckdb/duckdb-go-bindings/lib/linux-arm64 v0.10501.0/go.mod h1:81SGOYoEUs8qaAfSk1wRfM5oobrIJ5KI7AzYhK6/bvQ= +github.com/duckdb/duckdb-go-bindings/lib/windows-amd64 v0.10501.0 h1:jhhOonew2VOcTN4f+BOlnOTUgHEp797Uee2Tq8xMno8= +github.com/duckdb/duckdb-go-bindings/lib/windows-amd64 v0.10501.0/go.mod h1:K25pJL26ARblGDeuAkrdblFvUen92+CwksLtPEHRqqQ= +github.com/duckdb/duckdb-go/v2 v2.10501.0 h1:vYgvKBfotrZqpBESqHXYF5NVlbaYHRf0VrQEXtb/jnU= +github.com/duckdb/duckdb-go/v2 v2.10501.0/go.mod h1:825xmA19rJmdYWvSTd0kHWT9xq3EChSejO5RwevS9ZA= +github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI= +github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= +github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag= +github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE= +github.com/go-viper/mapstructure/v2 v2.5.0 h1:vM5IJoUAy3d7zRSVtIwQgBj7BiWtMPfmPEgAXnvj1Ro= +github.com/go-viper/mapstructure/v2 v2.5.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= +github.com/goccy/go-json v0.10.5 h1:Fq85nIqj+gXn/S5ahsiTlK3TmC85qgirsdTP/+DeaC4= +github.com/goccy/go-json v0.10.5/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M= +github.com/golang/freetype v0.0.0-20170609003504-e2365dfdc4a0 h1:DACJavvAHhabrF08vX0COfcOBJRhZ8lUbR+ZWIs0Y5g= +github.com/golang/freetype v0.0.0-20170609003504-e2365dfdc4a0/go.mod h1:E/TSTwGwJL78qG/PmXZO1EjYhfJinVAhrmmHX6Z8B9k= +github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek= +github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= +github.com/golang/snappy v1.0.0 h1:Oy607GVXHs7RtbggtPBnr2RmDArIsAefDwvrdWvRhGs= +github.com/golang/snappy v1.0.0/go.mod h1:/XxbfmMg8lxefKM7IXC3fBNl/7bRcc72aCRzEWrmP2Q= +github.com/google/flatbuffers v25.12.19+incompatible h1:haMV2JRRJCe1998HeW/p0X9UaMTK6SDo0ffLn2+DbLs= +github.com/google/flatbuffers v25.12.19+incompatible/go.mod h1:1AeVuKshWv4vARoZatz6mlQ0JxURH0Kv5+zNeJKJCa8= +github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= +github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/jonasknobloch/mbpe v0.1.1 h1:eXUrMdM7Wt6kPTA8cyX+t1PsInk3QIXpt430pE8CJEs= +github.com/jonasknobloch/mbpe v0.1.1/go.mod h1:2qW/5BfAu7GKAXXNYJzMrIPYk/8z5+/TbabaAb+vPv8= +github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= +github.com/klauspost/asmfmt v1.3.2 h1:4Ri7ox3EwapiOjCki+hw14RyKk201CN4rzyCJRFLpK4= +github.com/klauspost/asmfmt v1.3.2/go.mod h1:AG8TuvYojzulgDAMCnYn50l/5QV3Bs/tp6j0HLHbNSE= +github.com/klauspost/compress v1.18.3 h1:9PJRvfbmTabkOX8moIpXPbMMbYN60bWImDDU7L+/6zw= +github.com/klauspost/compress v1.18.3/go.mod h1:R0h/fSBs8DE4ENlcrlib3PsXS61voFxhIs2DeRhCvJ4= +github.com/klauspost/cpuid/v2 v2.3.0 h1:S4CRMLnYUhGeDFDqkGriYKdfoFlDnMtqTiI/sFzhA9Y= +github.com/klauspost/cpuid/v2 v2.3.0/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0= +github.com/minio/asm2plan9s v0.0.0-20200509001527-cdd76441f9d8 h1:AMFGa4R4MiIpspGNG7Z948v4n35fFGB3RR3G/ry4FWs= +github.com/minio/asm2plan9s v0.0.0-20200509001527-cdd76441f9d8/go.mod h1:mC1jAcsrzbxHt8iiaC+zU4b1ylILSosueou12R++wfY= +github.com/minio/c2goasm v0.0.0-20190812172519-36a3d3bbc4f3 h1:+n/aFZefKZp7spd8DFdX7uMikMLXX4oubIzJF4kv/wI= +github.com/minio/c2goasm v0.0.0-20190812172519-36a3d3bbc4f3/go.mod h1:RagcQ7I8IeTMnF8JTXieKnO4Z6JCsikNEzj0DwauVzE= +github.com/pierrec/lz4/v4 v4.1.25 h1:kocOqRffaIbU5djlIBr7Wh+cx82C0vtFb0fOurZHqD0= +github.com/pierrec/lz4/v4 v4.1.25/go.mod h1:EoQMVJgeeEOMsCqCzqFm2O0cJvljX2nGZjcRIPL34O4= +github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 h1:Jamvg5psRIccs7FGNTlIRMkT8wgtp5eCXdBlqhYGL6U= +github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY= +github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +github.com/xyproto/randomstring v1.0.5 h1:YtlWPoRdgMu3NZtP45drfy1GKoojuR7hmRcnhZqKjWU= +github.com/xyproto/randomstring v1.0.5/go.mod h1:rgmS5DeNXLivK7YprL0pY+lTuhNQW3iGxZ18UQApw/E= +github.com/yuin/goldmark v1.2.1/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74= +github.com/zeebo/assert v1.3.0 h1:g7C04CbJuIDKNPFHmsk4hwZDO5O+kntRxzaUoNXj+IQ= +github.com/zeebo/assert v1.3.0/go.mod h1:Pq9JiuJQpG8JLJdtkwrJESF0Foym2/D9XMU5ciN/wJ0= +github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs= +github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s= +go.jknobloc.com/x/dataset v0.0.0-20260410210408-456f8218a2d1 h1:W7C1aiXena5iyiUa4Oe+DVPcpfxgHhajOcALrEWUc+0= +go.jknobloc.com/x/dataset v0.0.0-20260410210408-456f8218a2d1/go.mod h1:UZypBoGqi23LRt29l2z9NtVD2pKBHd4unV/sztHRz4o= +go.jknobloc.com/x/llm v0.0.0-20260410210408-456f8218a2d1 h1:NNYSIsyYcmBBtaYFqCjXk1hIg0KskXTqyFbM9kwksfc= +go.jknobloc.com/x/llm v0.0.0-20260410210408-456f8218a2d1/go.mod h1:Ffq3FMFeZZyK3Xfx6eNYpBJZhzIBkytqmOEk+L7E6mE= +go.jknobloc.com/x/onnx v0.0.0-20260417134810-4d8f647bfb54 h1:KqukvTvgGj0cmE0uaIAyX5nA2vbgvcf6juGA4TMPL2U= +go.jknobloc.com/x/onnx v0.0.0-20260417134810-4d8f647bfb54/go.mod h1:eijjR4NKy4anm3jFnuuoZRDEFZ7CZgp8sOg+ZCWXF9c= +go.jknobloc.com/x/research/lesci v0.0.0-20260414175339-402477b200d1 h1:mLuKoovunuKq4eJYdCYgj+DRFtlV9AEWlL2iwoktjq4= +go.jknobloc.com/x/research/lesci v0.0.0-20260414175339-402477b200d1/go.mod h1:2xwgcfPnHyO07fVzRi0R3YLj6xj3dtmcjWxtyq7FMH4= +go.jknobloc.com/x/tensor v0.0.0-20260414175339-402477b200d1 h1:dauYEeyAjhzZ1hyyAR9exszAUKeFIqrrBNaAGboZaaA= +go.jknobloc.com/x/tensor v0.0.0-20260414175339-402477b200d1/go.mod h1:Citqq246efr+IGX0RBARqDm4gXwvPA4aNAz3F8Il6kk= +go.jknobloc.com/x/tokenizer v0.0.0-20260410210408-456f8218a2d1 h1:LmNp5OFgcm+4dsNrXPmMPUW/f0ZcvlwnRv5IzZAYY+8= +go.jknobloc.com/x/tokenizer v0.0.0-20260410210408-456f8218a2d1/go.mod h1:A6usfKV6fYSVHGvEMO7xMMIshky4NRCBzLhs36W7KC8= +go.jknobloc.com/x/tui v0.0.0-20260324194423-87bbece7e040 h1:7Ago/qaKuXIJ7/+/XDRO8Yb/4VWzBLKEcbAq2jD5k3Q= +go.jknobloc.com/x/tui v0.0.0-20260324194423-87bbece7e040/go.mod h1:kblmBsWO7WlRCo4sWpstaVV9ukNfYP3PeAPFhcoJvKA= +go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64= +go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y= +go.opentelemetry.io/otel v1.38.0 h1:RkfdswUDRimDg0m2Az18RKOsnI8UDzppJAtj01/Ymk8= +go.opentelemetry.io/otel v1.38.0/go.mod h1:zcmtmQ1+YmQM9wrNsTGV/q/uyusom3P8RxwExxkZhjM= +go.opentelemetry.io/otel/metric v1.38.0 h1:Kl6lzIYGAh5M159u9NgiRkmoMKjvbsKtYRwgfrA6WpA= +go.opentelemetry.io/otel/metric v1.38.0/go.mod h1:kB5n/QoRM8YwmUahxvI3bO34eVtQf2i4utNVLr9gEmI= +go.opentelemetry.io/otel/sdk v1.38.0 h1:l48sr5YbNf2hpCUj/FoGhW9yDkl+Ma+LrVl8qaM5b+E= +go.opentelemetry.io/otel/sdk v1.38.0/go.mod h1:ghmNdGlVemJI3+ZB5iDEuk4bWA3GkTpW+DOoZMYBVVg= +go.opentelemetry.io/otel/sdk/metric v1.38.0 h1:aSH66iL0aZqo//xXzQLYozmWrXxyFkBJ6qT5wthqPoM= +go.opentelemetry.io/otel/sdk/metric v1.38.0/go.mod h1:dg9PBnW9XdQ1Hd6ZnRz689CbtrUp0wMMs9iPcgT9EZA= +go.opentelemetry.io/otel/trace v1.38.0 h1:Fxk5bKrDZJUH+AMyyIXGcFAPah0oRcT+LuNtJrmcNLE= +go.opentelemetry.io/otel/trace v1.38.0/go.mod h1:j1P9ivuFsTceSWe1oY+EeW3sc+Pp42sO++GHkg4wwhs= +golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= +golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= +golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= +golang.org/x/exp v0.0.0-20260112195511-716be5621a96 h1:Z/6YuSHTLOHfNFdb8zVZomZr7cqNgTJvA8+Qz75D8gU= +golang.org/x/exp v0.0.0-20260112195511-716be5621a96/go.mod h1:nzimsREAkjBCIEFtHiYkrJyT+2uy9YZJB7H1k68CXZU= +golang.org/x/image v0.37.0 h1:ZiRjArKI8GwxZOoEtUfhrBtaCN+4b/7709dlT6SSnQA= +golang.org/x/image v0.37.0/go.mod h1:/3f6vaXC+6CEanU4KJxbcUZyEePbyKbaLoDOe4ehFYY= +golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= +golang.org/x/mod v0.33.0 h1:tHFzIWbBifEmbwtGz65eaWyGiGZatSrT9prnU8DbVL8= +golang.org/x/mod v0.33.0/go.mod h1:swjeQEj+6r7fODbD2cqrnje9PnziFuw4bmLbBZFrQ5w= +golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= +golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= +golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU= +golang.org/x/net v0.50.0 h1:ucWh9eiCGyDR3vtzso0WMQinm2Dnt8cFMuQa9K33J60= +golang.org/x/net v0.50.0/go.mod h1:UgoSli3F/pBgdJBHCTc+tp3gmrU4XswgGRgtnwWTfyM= +golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= +golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= +golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= +golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= +golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20210119212857-b64e53b001e4/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k= +golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= +golang.org/x/telemetry v0.0.0-20260209163413-e7419c687ee4 h1:bTLqdHv7xrGlFbvf5/TXNxy/iUwwdkjhqQTJDjW7aj0= +golang.org/x/telemetry v0.0.0-20260209163413-e7419c687ee4/go.mod h1:g5NllXBEermZrmR51cJDQxmJUHUOfRAaNyWBM+R+548= +golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= +golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= +golang.org/x/text v0.35.0 h1:JOVx6vVDFokkpaq1AEptVzLTpDe9KGpj5tR4/X+ybL8= +golang.org/x/text v0.35.0/go.mod h1:khi/HExzZJ2pGnjenulevKNX1W67CUy0AsXcNubPGCA= +golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= +golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= +golang.org/x/tools v0.1.0/go.mod h1:xkSsbof2nBLbhDlRMhhhyNLN/zl3eTqcnHD5viDpcZ0= +golang.org/x/tools v0.42.0 h1:uNgphsn75Tdz5Ji2q36v/nsFSfR/9BRFvqhGBaJGd5k= +golang.org/x/tools v0.42.0/go.mod h1:Ma6lCIwGZvHK6XtgbswSoWroEkhugApmsXyrUmBhfr0= +golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= +golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= +golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= +golang.org/x/xerrors v0.0.0-20240903120638-7835f813f4da h1:noIWHXmPHxILtqtCOPIhSt0ABwskkZKjD3bXGnZGpNY= +golang.org/x/xerrors v0.0.0-20240903120638-7835f813f4da/go.mod h1:NDW/Ps6MPRej6fsCIbMTohpP40sJ/P/vI1MoTEGwX90= +gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= +gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= +gonum.org/v1/plot v0.16.0 h1:dK28Qx/Ky4VmPUN/2zeW0ELyM6ucDnBAj5yun7M9n1g= +gonum.org/v1/plot v0.16.0/go.mod h1:Xz6U1yDMi6Ni6aaXILqmVIb6Vro8E+K7Q/GeeH+Pn0c= +google.golang.org/genproto/googleapis/rpc v0.0.0-20251029180050-ab9386a59fda h1:i/Q+bfisr7gq6feoJnS/DlpdwEL4ihp41fvRiM3Ork0= +google.golang.org/genproto/googleapis/rpc v0.0.0-20251029180050-ab9386a59fda/go.mod h1:7i2o+ce6H/6BluujYR+kqX3GKH+dChPTQU19wjRPiGk= +google.golang.org/grpc v1.78.0 h1:K1XZG/yGDJnzMdd/uZHAkVqJE+xIDOcmdSFZkBUicNc= +google.golang.org/grpc v1.78.0/go.mod h1:I47qjTo4OKbMkjA/aOOwxDIiPSBofUtQUI5EfpWvW7U= +google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= +google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +honnef.co/go/tools v0.1.3/go.mod h1:NgwopIslSNH47DimFoV78dnkksY2EFtX0ajyb3K/las= +rsc.io/pdf v0.1.1 h1:k1MczvYDUvJBe93bYd7wrZLLUEcLZAuF824/I4e5Xr4= +rsc.io/pdf v0.1.1/go.mod h1:n8OzWcQ6Sp37PL01nO98y4iUCRdTGarVfzxY20ICaU4= diff --git a/research/sander/plot.go b/research/sander/plot.go new file mode 100644 index 0000000..bc6c483 --- /dev/null +++ b/research/sander/plot.go @@ -0,0 +1,100 @@ +package sander + +import ( + "context" + "database/sql" + "image/color" + "path" + + "gonum.org/v1/plot" + "gonum.org/v1/plot/plotter" + "gonum.org/v1/plot/vg" + "gonum.org/v1/plot/vg/draw" +) + +func (e *Experiment) Plot(db *sql.DB) error { + var rows *sql.Rows + + if r, err := db.QueryContext(context.Background(), ` + SELECT s.token_id, s.distance, e.reference + FROM similarity_buffer s + JOIN embeddings e ON s.token_id = e.token_id + ORDER BY s.token_id ASC + `); err != nil { + return err + } else { + rows = r + } + + defer rows.Close() + + var pts plotter.XYZs + + for rows.Next() { + var id, distance float64 + var reference bool + + if err := rows.Scan(&id, &distance, &reference); err != nil { + return err + } + + z := 0.0 + + if reference { + z = 1.0 + } + + pts = append(pts, plotter.XYZ{X: id, Y: distance, Z: z}) + } + + if err := rows.Err(); err != nil { + return err + } + + return render(pts, path.Join(e.name, "similarity_scatter.png")) +} + +func render(pts plotter.XYZs, out string) error { + p := plot.New() + + p.Y.Label.Text = "Euclidean Distance" + + p.Y.Min = 0 + p.Y.Max = 5 + + var candidate, reference plotter.XYZs + + for _, pt := range pts { + if pt.Z == 1.0 { + reference = append(reference, pt) + } else { + candidate = append(candidate, pt) + } + } + + if err := addScatter(p, candidate, color.RGBA{A: 64}); err != nil { + return err + } + + if err := addScatter(p, reference, color.RGBA{R: 255, G: 69, B: 0, A: 255}); err != nil { + return err + } + + return p.Save(10*vg.Inch, 6*vg.Inch, out) +} + +func addScatter(p *plot.Plot, pts plotter.XYZs, c color.Color) error { + s, err := plotter.NewScatter(pts) + + if err != nil { + return err + } + + s.Color = c + s.Radius = vg.Points(2) + s.Shape = draw.CircleGlyph{} + + p.Add(s) + + return nil +} diff --git a/research/sander/reference.go b/research/sander/reference.go new file mode 100644 index 0000000..d6aa466 --- /dev/null +++ b/research/sander/reference.go @@ -0,0 +1,59 @@ +package sander + +import ( + "go.jknobloc.com/x/tokenizer/bpe" +) + +func UnusedTokensGPT2(tokenizer *bpe.Tokenizer) map[int64]struct{} { + unused := map[int64]struct{}{ + 177: {}, + 178: {}, + 179: {}, + 180: {}, + 181: {}, + 182: {}, + 183: {}, + 184: {}, + 185: {}, + 186: {}, + 187: {}, + } + + alphabet := bpe.InitialAlphabet() + + itoa := bpe.Itoa(tokenizer) + + for id := range unused { + token, ok := itoa[id] + + if !ok { + panic("unknown token id") + } + + if token != string(alphabet[id]) { + panic("unexpected token") + } + } + + return unused +} + +func UnusedTokensMBPE(tokenizer *bpe.Tokenizer) map[int64]struct{} { + unused := make(map[int64]struct{}) + + vocab := bpe.Vocab(tokenizer) + + mask := bpe.ReachableTokens(tokenizer, vocab) + + for i := range vocab { + if mask[i] { + continue + } + + unused[int64(i)] = struct{}{} + } + + return unused +} + +// TODO we could just filter some input data |
