From 4386784c1006321c5cf0e1e005b19e0a7d7563ae Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Wed, 8 Apr 2026 22:24:10 +0200 Subject: Refactor model options * Add options struct --- gpt2/model.go | 18 +++++++----------- 1 file changed, 7 insertions(+), 11 deletions(-) (limited to 'gpt2/model.go') diff --git a/gpt2/model.go b/gpt2/model.go index 0708c7e..20dcb71 100644 --- a/gpt2/model.go +++ b/gpt2/model.go @@ -15,21 +15,17 @@ type Model struct { name string deviceID string config Config - withCache bool - withLogits bool - withLogProbs bool + options Options session *ort.DynamicAdvancedSession allocator *Allocator } -func NewModel(name string, deviceID string, cfg Config, withCache bool, withLogits bool, withLogProbs bool) *Model { +func NewModel(name string, deviceID string, cfg Config, opts Options) *Model { return &Model{ name: name, deviceID: deviceID, config: cfg, - withCache: withCache, - withLogits: withLogits, - withLogProbs: withLogProbs, + options: opts, } } @@ -60,7 +56,7 @@ func IntraOpNumThreads() int { } func (m *Model) Init() error { - m.allocator = NewAllocator(m.config, 1, m.withCache, m.withLogits, m.withLogProbs) + m.allocator = NewAllocator(m.config, m.options, 1) var options *ort.SessionOptions @@ -98,11 +94,11 @@ func (m *Model) Destroy() { } func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]int64, error) { - if !m.withLogits { + if !m.options.WithLogits { panic("generate requires logits output") } - if steps > 0 && !m.withCache { + if steps > 0 && !m.options.WithCache { panic("generate with steps > 0 requires cache") } @@ -154,7 +150,7 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in } func (m *Model) Score(tokens []int64, batchSize int, logProbs *[]float32) error { - if !m.withLogProbs { + if !m.options.WithLogProbs { panic("score requires token_logprobs output") } -- cgit v1.3.1