summaryrefslogtreecommitdiff
path: root/llm/run.go
diff options
context:
space:
mode:
Diffstat (limited to 'llm/run.go')
-rw-r--r--llm/run.go30
1 files changed, 17 insertions, 13 deletions
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
+ }
}