diff options
Diffstat (limited to 'llm')
| -rw-r--r-- | llm/cmd/eval/logprobs.go | 32 | ||||
| -rw-r--r-- | llm/cmd/eval/ppl.go | 23 | ||||
| -rw-r--r-- | llm/evaluator.go | 28 |
3 files changed, 39 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) diff --git a/llm/evaluator.go b/llm/evaluator.go index 0eb5418..8f2a295 100644 --- a/llm/evaluator.go +++ b/llm/evaluator.go @@ -1,7 +1,11 @@ package llm import ( + "errors" + "sync" "sync/atomic" + + "go.jknobloc.com/x/dataset" ) type Evaluator[R any] struct { @@ -35,3 +39,27 @@ func NewEvaluator[R any](model Causal, tokenizer Tokenizer, callback func(job Jo func (e *Evaluator[R]) Results() chan R { return e.results } + +func (e *Evaluator[R]) RunAndCollect(data dataset.Reader, window, stride int, callback func(R) error) error { + var wg sync.WaitGroup + + var collectErr error + + wg.Add(1) + + go func() { + defer wg.Done() + + for r := range e.results { + if err := callback(r); err != nil && collectErr == nil { + collectErr = err + } + } + }() + + runErr := e.Run(data, window, stride) + + wg.Wait() + + return errors.Join(runErr, collectErr) +} |
