From 59b0d9b9e9e648b3626cc7a17cfc769e6bfa6a9c Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Tue, 7 Apr 2026 19:45:39 +0200 Subject: Batch scoring --- llm/run.go | 30 +++++++++++++++++------------- 1 file changed, 17 insertions(+), 13 deletions(-) (limited to 'llm/run.go') diff --git a/llm/run.go b/llm/run.go index 9cdf8c3..21f60eb 100644 --- a/llm/run.go +++ b/llm/run.go @@ -113,30 +113,34 @@ func (e *Evaluator[R]) Run(title string, data dataset.Reader, window, stride int } func (e *Evaluator[R]) execute(j *batch, device int) { - if j.Size() != 1 { - panic("unimplemented") - } + b := j.Size() + + s := len(j.jobs[0].Tokens) - job := j.jobs[0] + tokens := make([]int64, b*s) - if job.Seen < 1 { - panic("empty context") + for i, job := range j.jobs { + copy(tokens[i*s:], job.Tokens) } m := e.models[device] - logProbs := make([]float32, 0, len(job.Tokens)-1) + logProbs := make([]float32, 0, b*(s-1)) - if err := m.Score(job.Tokens, &logProbs); err != nil { + if err := m.Score(tokens, b, &logProbs); err != nil { panic(err) // TODO handle } - l := logProbs[job.Seen-1:] - t := toInt(job.Tokens[job.Seen:]) + for i, job := range j.jobs { + if job.Seen < 1 { + panic("empty context") + } - r := e.callback(job, l, t) + l := logProbs[i*(s-1)+job.Seen-1 : (i+1)*(s-1)] + t := toInt(job.Tokens[job.Seen:]) - e.results <- r + r := e.callback(job, l, t) - return + e.results <- r + } } -- cgit v1.3.1