summaryrefslogtreecommitdiff
path: root/llm/cmd/eval/logprobs.go
blob: 7228ea6639e318fe485b23f0888926de7a204994 (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
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
package main

import (
	"context"
	"database/sql"
	"log"

	_ "github.com/duckdb/duckdb-go/v2"

	"go.jknobloc.com/x/llm"
)

type logProb struct {
	document int
	token    int
	value    float32
	offset   int
}

type logProbs []logProb

func logprobs() {
	d := data()

	m := model()
	t := tokenizer()

	e := llm.NewEvaluator(m, t, func(j llm.Job, l []float32, tokens []int) logProbs {
		r := make(logProbs, len(tokens))

		for i, token := range tokens {
			r[i] = logProb{
				document: j.Document,
				token:    token,
				value:    l[i],
				offset:   i,
			}
		}

		return r
	})

	var insertStmt *sql.Stmt

	if stmt, db, err := prepare("logprobs.db?access_mode=READ_WRITE"); err != nil {
		log.Fatal(err)
	} else {
		insertStmt = stmt

		defer db.Close()
		defer stmt.Close()
	}

	if err := e.RunAndCollect("LogProbs", d, 1024, 512, func(r logProbs) error {
		for _, l := range r {
			if err := insert(insertStmt, l); err != nil {
				return err
			}
		}

		return nil
	}); err != nil {
		log.Fatal(err)
	}
}

func prepare(name string) (*sql.Stmt, *sql.DB, error) {
	var db *sql.DB

	if database, err := sql.Open("duckdb", name); err != nil {
		return nil, nil, err
	} else {
		db = database
	}

	if err := db.Ping(); err != nil {
		_ = db.Close()

		return nil, nil, err
	}

	queryCreate := "CREATE TABLE context(uid INTEGER, token INTEGER, logProb FLOAT, pos INTEGER)"

	if _, err := db.ExecContext(context.Background(), queryCreate); err != nil {
		_ = db.Close()

		return nil, nil, err
	}

	var stmt *sql.Stmt

	queryInsert := "INSERT INTO context VALUES(?, ?, ?, ?)"

	if statement, err := db.PrepareContext(context.Background(), queryInsert); err != nil {
		_ = db.Close()

		return nil, nil, err
	} else {
		stmt = statement
	}

	return stmt, db, nil
}

func insert(stmt *sql.Stmt, prob logProb) error {
	if _, err := stmt.ExecContext(context.Background(), prob.document, prob.token, prob.value, prob.offset); err != nil {
		return err
	}

	return nil
}