summaryrefslogtreecommitdiff
path: root/research/sander/analyze.go
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-17 16:07:14 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-17 16:07:14 +0200
commit3fe37224ccc2500e7e767e5fc00163ef61182828 (patch)
tree7bd14e080c080b24026660beda42817a3c9dc77f /research/sander/analyze.go
parent4d8f647bfb54899e3052931790e288da7af97050 (diff)
Add sander module
Diffstat (limited to 'research/sander/analyze.go')
-rw-r--r--research/sander/analyze.go49
1 files changed, 49 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()
+}