summaryrefslogtreecommitdiff
path: root/llm/cmd/eval
diff options
context:
space:
mode:
Diffstat (limited to 'llm/cmd/eval')
-rw-r--r--llm/cmd/eval/logprobs.go177
-rw-r--r--llm/cmd/eval/main.go1
2 files changed, 178 insertions, 0 deletions
diff --git a/llm/cmd/eval/logprobs.go b/llm/cmd/eval/logprobs.go
new file mode 100644
index 0000000..eaa7459
--- /dev/null
+++ b/llm/cmd/eval/logprobs.go
@@ -0,0 +1,177 @@
+package main
+
+import (
+ "context"
+ "database/sql"
+ "log"
+ "math"
+ "slices"
+ "sync"
+
+ _ "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[logProbs]()
+
+ e.SetTokenizer(t)
+ e.AddModel(m)
+
+ e.SetCallback(func(j llm.Job, logits [][]float32, tokens []int) logProbs {
+ probs := selectLogProbs(logits, tokens)
+
+ r := make(logProbs, len(tokens))
+
+ for i, token := range tokens {
+ r[i] = logProb{
+ document: j.Document,
+ token: token,
+ value: probs[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()
+ }
+
+ results := e.Results()
+
+ var wg sync.WaitGroup
+ var writeErr error
+
+ wg.Add(1)
+
+ go func() {
+ defer wg.Done()
+
+ for r := range results {
+ for _, l := range r {
+ if err := insert(insertStmt, l); err != nil {
+ writeErr = err
+
+ return
+ }
+ }
+ }
+ }()
+
+ if err := e.Run(d, 1024, 512); err != nil {
+ log.Fatal(err)
+ }
+
+ wg.Wait()
+
+ if writeErr != nil {
+ log.Fatal(writeErr)
+ }
+}
+
+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
+}
+
+func selectLogProbs(logits [][]float32, tokens []int) []float32 {
+ if len(logits) != len(tokens) {
+ panic("length mismatch")
+ }
+
+ r := make([]float32, len(tokens))
+
+ for i, token := range tokens {
+ logprobs := logSoftmax(logits[i])
+
+ r[i] = logprobs[token]
+ }
+
+ return r
+}
+
+func logSoftmax(logits []float32) []float32 {
+ m := slices.Max(logits)
+
+ s := float32(0.0)
+ r := make([]float32, len(logits))
+
+ for i, v := range logits {
+ e := float32(math.Exp(float64(v - m)))
+
+ r[i] = v
+ s += e
+ }
+
+ lse := float32(math.Log(float64(s))) + m
+
+ for i := range r {
+ r[i] -= lse
+ }
+
+ return r
+}
diff --git a/llm/cmd/eval/main.go b/llm/cmd/eval/main.go
index 05b3f96..758adad 100644
--- a/llm/cmd/eval/main.go
+++ b/llm/cmd/eval/main.go
@@ -11,6 +11,7 @@ import (
func main() {
perplexity()
+ // logprobs()
}
func data() *dataset.ParquetReader {