summaryrefslogtreecommitdiff
path: root/research/sander/analyze.go
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()
}