From 408536516bdfedc93d93226dc08280ed224361fc Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Fri, 24 Apr 2026 11:14:43 +0200 Subject: Cleanup eval command * Remove logprob extraction --- llm/cmd/eval/logprobs.go | 118 --------------------------------------------- llm/cmd/eval/main.go | 1 - llm/cmd/eval/perplexity.go | 62 ++++++++++++++++++++++++ llm/cmd/eval/ppl.go | 78 ------------------------------ 4 files changed, 62 insertions(+), 197 deletions(-) delete mode 100644 llm/cmd/eval/logprobs.go create mode 100644 llm/cmd/eval/perplexity.go delete mode 100644 llm/cmd/eval/ppl.go diff --git a/llm/cmd/eval/logprobs.go b/llm/cmd/eval/logprobs.go deleted file mode 100644 index f39e351..0000000 --- a/llm/cmd/eval/logprobs.go +++ /dev/null @@ -1,118 +0,0 @@ -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 - }, llm.EvaluatorConfig{ - BatchSize: 1, - NumWorkers: 4, - }) - - 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) - } - - if err := m.Destroy(); 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 -} diff --git a/llm/cmd/eval/main.go b/llm/cmd/eval/main.go index 4d27ada..d31bc32 100644 --- a/llm/cmd/eval/main.go +++ b/llm/cmd/eval/main.go @@ -15,7 +15,6 @@ func main() { } perplexity() - // logprobs() if err := gpt2.DestroyEnvironment(); err != nil { log.Fatal(err) diff --git a/llm/cmd/eval/perplexity.go b/llm/cmd/eval/perplexity.go new file mode 100644 index 0000000..ba92c99 --- /dev/null +++ b/llm/cmd/eval/perplexity.go @@ -0,0 +1,62 @@ +package main + +import ( + "fmt" + "log" + "math" + + "go.jknobloc.com/x/llm" +) + +type pplResult struct { + v float64 + n int +} + +func perplexity() { + d := data() + + m := model() + t := tokenizer() + + e := llm.NewEvaluator(m, t, func(job llm.Job, logProbs []float32, tokens []int) pplResult { + total := float64(0) + + n := 0 + + for _, p := range logProbs { + total -= float64(p) + + n++ + } + + return pplResult{ + v: total, + n: n, + } + }, llm.EvaluatorConfig{ + BatchSize: 1, + NumWorkers: 4, + }) + + total := float64(0) + n := 0 + + if err := e.RunAndCollect("Perplexity", d, 1024, 512, func(r pplResult) error { + total += r.v + n += r.n + + return nil + }); err != nil { + log.Fatal(err) + } + + avg := total / float64(n) + ppl := math.Exp(avg) + + fmt.Println(ppl) + + if err := m.Destroy(); err != nil { + log.Fatal(err) + } +} diff --git a/llm/cmd/eval/ppl.go b/llm/cmd/eval/ppl.go deleted file mode 100644 index 2394d83..0000000 --- a/llm/cmd/eval/ppl.go +++ /dev/null @@ -1,78 +0,0 @@ -package main - -import ( - "fmt" - "log" - "math" - "strings" - - "go.jknobloc.com/x/dataset" - "go.jknobloc.com/x/llm" -) - -type pplResult struct { - v float64 - n int -} - -func perplexity() { - d := data() - - m := model() - t := tokenizer() - - e := llm.NewEvaluator(m, t, func(job llm.Job, logProbs []float32, tokens []int) pplResult { - total := float64(0) - - n := 0 - - for _, p := range logProbs { - total -= float64(p) - - n++ - } - - return pplResult{ - v: total, - n: n, - } - }, llm.EvaluatorConfig{ - BatchSize: 1, - NumWorkers: 4, - }) - - total := float64(0) - n := 0 - - if err := e.RunAndCollect("Perplexity", d, 1024, 512, func(r pplResult) error { - total += r.v - n += r.n - - return nil - }); err != nil { - log.Fatal(err) - } - - avg := total / float64(n) - ppl := math.Exp(avg) - - fmt.Println(ppl) - - if err := m.Destroy(); err != nil { - log.Fatal(err) - } -} - -func joined() dataset.Reader { - miniPile := data() - - docs := make([]string, 0) - - for _, d := range miniPile.Texts() { - docs = append(docs, d) - } - - j := dataset.NewStringReader(strings.Join(docs, "\n\n")) - - return j -} -- cgit v1.2.3