summaryrefslogtreecommitdiff
path: root/gpt2/model.go
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-06 22:39:07 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-06 23:27:30 +0200
commita8f1b564dba5cc4029ff25a49d2aa47bc7dab406 (patch)
treee601fd866015175c740699aa1033a6e8d708eec8 /gpt2/model.go
parent42ee72bf62a03cffc0d10bb691a1b011dfadccfa (diff)
Add parameter to control logits output binding
Diffstat (limited to 'gpt2/model.go')
-rw-r--r--gpt2/model.go12
1 files changed, 7 insertions, 5 deletions
diff --git a/gpt2/model.go b/gpt2/model.go
index 214affe..19a3363 100644
--- a/gpt2/model.go
+++ b/gpt2/model.go
@@ -16,17 +16,19 @@ type Model struct {
deviceID string
config Config
withCache bool
+ withLogits bool
withLogProbs bool
session *ort.DynamicAdvancedSession
allocator *Allocator
}
-func NewModel(name string, deviceID string, config Config, withCache bool, withLogProbs bool) *Model {
+func NewModel(name string, deviceID string, config Config, withCache bool, withLogits bool, withLogProbs bool) *Model {
return &Model{
name: name,
deviceID: deviceID,
config: config,
withCache: withCache,
+ withLogits: withLogits,
withLogProbs: withLogProbs,
}
}
@@ -64,7 +66,7 @@ func (m *Model) Init() error {
return err
}
- m.allocator = NewAllocator(m.config, m.withCache, m.withLogProbs)
+ m.allocator = NewAllocator(m.config, m.withCache, m.withLogits, m.withLogProbs)
var options *ort.SessionOptions
@@ -104,8 +106,8 @@ func (m *Model) Destroy() error {
}
func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]int64, error) {
- if m.withLogProbs {
- panic("generate called on eval model")
+ if !m.withLogits {
+ panic("generate requires logits output")
}
if len(prompt) == 0 {
@@ -157,7 +159,7 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in
func (m *Model) Score(tokens []int64, logProbs *[]float32) error {
if !m.withLogProbs {
- panic("score called on default model")
+ panic("score requires log_probs output")
}
if err := m.allocator.Init(tokens); err != nil {