summaryrefslogtreecommitdiff
path: root/gpt2/model.go
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 /gpt2/model.go
parent96b4edc5e291e253747956083933339040ad07c3 (diff)
Batch scoring
Diffstat (limited to 'gpt2/model.go')
-rw-r--r--gpt2/model.go6
1 files changed, 4 insertions, 2 deletions
diff --git a/gpt2/model.go b/gpt2/model.go
index 2b23cdf..a5f4d27 100644
--- a/gpt2/model.go
+++ b/gpt2/model.go
@@ -60,7 +60,7 @@ func IntraOpNumThreads() int {
}
func (m *Model) Init() error {
- m.allocator = NewAllocator(m.config, m.withCache, m.withLogits, m.withLogProbs)
+ m.allocator = NewAllocator(m.config, 1, m.withCache, m.withLogits, m.withLogProbs)
var options *ort.SessionOptions
@@ -153,11 +153,13 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in
return r, nil
}
-func (m *Model) Score(tokens []int64, logProbs *[]float32) error {
+func (m *Model) Score(tokens []int64, batchSize int, logProbs *[]float32) error {
if !m.withLogProbs {
panic("score requires token_logprobs output")
}
+ m.allocator.SetBatchSize(batchSize)
+
if err := m.allocator.Init(tokens); err != nil {
return err
}