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/evaluator.go | |
| parent | a53c085de9f3e27dadf1f3f3025708d5c7542f92 (diff) | |
Add helper to collect results
Diffstat (limited to 'llm/evaluator.go')
| -rw-r--r-- | llm/evaluator.go | 28 |
1 files changed, 28 insertions, 0 deletions
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) +} |
