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 /gpt2/model.go | |
| parent | 96b4edc5e291e253747956083933339040ad07c3 (diff) | |
Batch scoring
Diffstat (limited to 'gpt2/model.go')
| -rw-r--r-- | gpt2/model.go | 6 |
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 } |
