summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-05-04 15:39:45 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-05-04 15:39:45 +0200
commit29c52bf0f0981d4ea8e12e86237365da2b225ef7 (patch)
tree6ac95a78f2d0ad59fbafa760c216835f38f462e1
parent281e82eed24c1f871e868afbd23cb6811333c2f4 (diff)
Add parameter to set hidden dimension
-rw-r--r--research/sander/analyze.go10
-rw-r--r--research/sander/cmd/sander/main.go3
-rw-r--r--research/sander/experiment.go4
-rw-r--r--research/sander/extract.go11
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
}