summaryrefslogtreecommitdiff
path: root/research/sander/analyze.go
diff options
context:
space:
mode:
Diffstat (limited to 'research/sander/analyze.go')
-rw-r--r--research/sander/analyze.go10
1 files changed, 5 insertions, 5 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)
}