blob: 6abc9b8885a21854332bbd7dcfcac82466519ed3 (
plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
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[%d] 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
}
query := fmt.Sprintf(sqlSimilarity, e.hiddenDim)
if _, err := tx.ExecContext(ctx, query); err != nil {
return fmt.Errorf("similarity: %w", err)
}
return tx.Commit()
}
|