diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-03-25 21:38:10 +0100 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-03-25 21:44:51 +0100 |
| commit | 38e6a5a73a901f9086b10debb938f672ca87a3b0 (patch) | |
| tree | dd79e0452ceadcf408092056b82f9cf2383f9691 /llm/cmd/eval | |
| parent | a53c085de9f3e27dadf1f3f3025708d5c7542f92 (diff) | |
Add helper to collect results
Diffstat (limited to 'llm/cmd/eval')
| -rw-r--r-- | llm/cmd/eval/logprobs.go | 32 | ||||
| -rw-r--r-- | llm/cmd/eval/ppl.go | 23 |
2 files changed, 11 insertions, 44 deletions
diff --git a/llm/cmd/eval/logprobs.go b/llm/cmd/eval/logprobs.go index 7c04626..264f943 100644 --- a/llm/cmd/eval/logprobs.go +++ b/llm/cmd/eval/logprobs.go @@ -6,7 +6,6 @@ import ( "log" "math" "slices" - "sync" _ "github.com/duckdb/duckdb-go/v2" @@ -56,36 +55,17 @@ func logprobs() { 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.RunAndCollect(d, 1024, 512, func(r logProbs) error { + for _, l := range r { + if err := insert(insertStmt, l); err != nil { + return err } } - }() - if err := e.Run(d, 1024, 512); err != nil { + return nil + }); err != nil { log.Fatal(err) } - - wg.Wait() - - if writeErr != nil { - log.Fatal(writeErr) - } } func prepare(name string) (*sql.Stmt, *sql.DB, error) { diff --git a/llm/cmd/eval/ppl.go b/llm/cmd/eval/ppl.go index 0a028d2..c03c67c 100644 --- a/llm/cmd/eval/ppl.go +++ b/llm/cmd/eval/ppl.go @@ -5,7 +5,6 @@ import ( "log" "math" "strings" - "sync" "go.jknobloc.com/x/dataset" "go.jknobloc.com/x/llm" @@ -31,30 +30,18 @@ func perplexity() { } }) - results := e.Results() - - var wg sync.WaitGroup - total := float64(0) n := 0 - wg.Add(1) - - go func() { - defer wg.Done() + if err := e.RunAndCollect(d, 1024, 512, func(r pplResult) error { + total += r.v + n += r.n - for r := range results { - total += r.v - n += r.n - } - }() - - if err := e.Run(d, 1024, 512); err != nil { + return nil + }); err != nil { log.Fatal(err) } - wg.Wait() - avg := total / float64(n) ppl := math.Exp(avg) |
