summaryrefslogtreecommitdiff
path: root/llm
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-07 19:45:39 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-07 19:45:39 +0200
commit59b0d9b9e9e648b3626cc7a17cfc769e6bfa6a9c (patch)
tree3aa0d434d501a5f2c0b08790fc982f998f88c937 /llm
parent96b4edc5e291e253747956083933339040ad07c3 (diff)
Batch scoring
Diffstat (limited to 'llm')
-rw-r--r--llm/causal.go2
-rw-r--r--llm/run.go30
2 files changed, 18 insertions, 14 deletions
diff --git a/llm/causal.go b/llm/causal.go
index b683985..69c3127 100644
--- a/llm/causal.go
+++ b/llm/causal.go
@@ -2,5 +2,5 @@ package llm
type Causal interface {
Generate(prompt []int64, steps int64, logits *[][]float32) ([]int64, error)
- Score(tokens []int64, logProbs *[]float32) error
+ Score(tokens []int64, batchSize int, logProbs *[]float32) error
}
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
+ }
}