diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-07 19:45:39 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-07 19:45:39 +0200 |
| commit | 59b0d9b9e9e648b3626cc7a17cfc769e6bfa6a9c (patch) | |
| tree | 3aa0d434d501a5f2c0b08790fc982f998f88c937 /llm/run.go | |
| parent | 96b4edc5e291e253747956083933339040ad07c3 (diff) | |
Batch scoring
Diffstat (limited to 'llm/run.go')
| -rw-r--r-- | llm/run.go | 30 |
1 files changed, 17 insertions, 13 deletions
@@ -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 + } } |
