summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--llm/cmd/eval/logprobs.go118
-rw-r--r--llm/cmd/eval/main.go1
-rw-r--r--llm/cmd/eval/perplexity.go (renamed from llm/cmd/eval/ppl.go)16
3 files changed, 0 insertions, 135 deletions
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/ppl.go b/llm/cmd/eval/perplexity.go
index 2394d83..ba92c99 100644
--- a/llm/cmd/eval/ppl.go
+++ b/llm/cmd/eval/perplexity.go
@@ -4,9 +4,7 @@ import (
"fmt"
"log"
"math"
- "strings"
- "go.jknobloc.com/x/dataset"
"go.jknobloc.com/x/llm"
)
@@ -62,17 +60,3 @@ func perplexity() {
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
-}