summaryrefslogtreecommitdiff
path: root/llm/cmd
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-03-25 21:38:10 +0100
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-03-25 21:44:51 +0100
commit38e6a5a73a901f9086b10debb938f672ca87a3b0 (patch)
treedd79e0452ceadcf408092056b82f9cf2383f9691 /llm/cmd
parenta53c085de9f3e27dadf1f3f3025708d5c7542f92 (diff)
Add helper to collect results
Diffstat (limited to 'llm/cmd')
-rw-r--r--llm/cmd/eval/logprobs.go32
-rw-r--r--llm/cmd/eval/ppl.go23
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)