summaryrefslogtreecommitdiff
path: root/llm/cmd/eval/logprobs.go
diff options
context:
space:
mode:
Diffstat (limited to 'llm/cmd/eval/logprobs.go')
-rw-r--r--llm/cmd/eval/logprobs.go32
1 files changed, 6 insertions, 26 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) {