From 29c52bf0f0981d4ea8e12e86237365da2b225ef7 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Mon, 4 May 2026 15:39:45 +0200 Subject: Add parameter to set hidden dimension --- research/sander/analyze.go | 10 +++++----- research/sander/cmd/sander/main.go | 3 ++- research/sander/experiment.go | 4 +++- research/sander/extract.go | 11 ++++++----- 4 files changed, 16 insertions(+), 12 deletions(-) diff --git a/research/sander/analyze.go b/research/sander/analyze.go index c58195d..6abc9b8 100644 --- a/research/sander/analyze.go +++ b/research/sander/analyze.go @@ -6,11 +6,10 @@ import ( "fmt" ) -var ( - sqlSimilarity = ` +var sqlSimilarity = ` CREATE TABLE similarity_buffer AS WITH centroid AS ( - SELECT list(avg_val ORDER BY idx)::FLOAT[768] AS vec + SELECT list(avg_val ORDER BY idx)::FLOAT[%d] AS vec FROM ( SELECT idx, avg(val) AS avg_val FROM embeddings @@ -22,7 +21,6 @@ var ( 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() @@ -41,7 +39,9 @@ func (e *Experiment) Analyze(db *sql.DB) error { return err } - if _, err := tx.ExecContext(ctx, sqlSimilarity); err != nil { + query := fmt.Sprintf(sqlSimilarity, e.hiddenDim) + + if _, err := tx.ExecContext(ctx, query); err != nil { return fmt.Errorf("similarity: %w", err) } diff --git a/research/sander/cmd/sander/main.go b/research/sander/cmd/sander/main.go index 56f1488..ae7f7b9 100644 --- a/research/sander/cmd/sander/main.go +++ b/research/sander/cmd/sander/main.go @@ -5,6 +5,7 @@ import ( "log" "path" + "go.jknobloc.com/x/gpt2" "go.jknobloc.com/x/research/sander" "go.jknobloc.com/x/shelf" "go.jknobloc.com/x/tokenizer/bpe" @@ -33,7 +34,7 @@ func run(src, dst string) error { } t := must(bpe.NewTokenizerFromFiles(path.Join(src, "vocab.json"), path.Join(src, "merges.txt"), cfg)) - e := must(sander.NewExperiment(dst, path.Join(src, "model.onnx"), sander.UnusedTokensMBPE(t))) + e := must(sander.NewExperiment(dst, path.Join(src, "model.onnx"), sander.UnusedTokensMBPE(t), gpt2.ConfigDefault().HiddenDim)) return e.Run() } diff --git a/research/sander/experiment.go b/research/sander/experiment.go index 1521cf8..bcad99c 100644 --- a/research/sander/experiment.go +++ b/research/sander/experiment.go @@ -14,13 +14,15 @@ type Experiment struct { name string model string reference map[int64]struct{} + hiddenDim int } -func NewExperiment(name, model string, reference map[int64]struct{}) (*Experiment, error) { +func NewExperiment(name, model string, reference map[int64]struct{}, hiddenDim int) (*Experiment, error) { e := &Experiment{ name: name, model: model, reference: reference, + hiddenDim: hiddenDim, } return e, nil diff --git a/research/sander/extract.go b/research/sander/extract.go index de31b1f..c72fe7f 100644 --- a/research/sander/extract.go +++ b/research/sander/extract.go @@ -3,6 +3,7 @@ package sander import ( "database/sql" "database/sql/driver" + "fmt" "go.jknobloc.com/x/onnx" "go.jknobloc.com/x/research/lesci" @@ -20,16 +21,16 @@ func (e *Experiment) Extract(db *sql.DB) error { shape = s } - if shape[1] != 768 { - panic("unimplemented") + if shape[1] != e.hiddenDim { + panic("shape mismatch") } t := tensor.NewDense[float32](shape, data) - if _, err := db.Exec(` + if _, err := db.Exec(fmt.Sprintf(` DROP TABLE IF EXISTS embeddings; - CREATE TABLE embeddings (token_id INTEGER, embedding FLOAT[768], reference BOOLEAN); - `); err != nil { + CREATE TABLE embeddings (token_id INTEGER, embedding FLOAT[%d], reference BOOLEAN); + `, e.hiddenDim)); err != nil { return err } -- cgit v1.2.3