diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-08 22:24:10 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-08 23:46:10 +0200 |
| commit | 4386784c1006321c5cf0e1e005b19e0a7d7563ae (patch) | |
| tree | 3d22360aaac8e4d64ca90e39b7247942073582d1 | |
| parent | a006507a24b93a6af7e3cff8833fb96b48b159d7 (diff) | |
Refactor model options
* Add options struct
| -rw-r--r-- | gpt2/allocator.go | 46 | ||||
| -rw-r--r-- | gpt2/cmd/gpt2/main.go | 16 | ||||
| -rw-r--r-- | gpt2/config.go | 6 | ||||
| -rw-r--r-- | gpt2/model.go | 18 | ||||
| -rw-r--r-- | gpt2/model_test.go | 8 | ||||
| -rw-r--r-- | llm/cmd/eval/main.go | 8 |
6 files changed, 62 insertions, 40 deletions
diff --git a/gpt2/allocator.go b/gpt2/allocator.go index 3c85410..46d612d 100644 --- a/gpt2/allocator.go +++ b/gpt2/allocator.go @@ -8,32 +8,28 @@ import ( type Allocator struct { config Config + options Options batchSize int sequenceLength int64 step int64 inputNames []string outputNames []string values map[string]ort.Value - withCache bool - withLogits bool - withLogProbs bool } -func NewAllocator(cfg Config, batchSize int, withCache bool, withLogits bool, withLogProbs bool) *Allocator { +func NewAllocator(cfg Config, opts Options, batchSize int) *Allocator { return &Allocator{ config: cfg, + options: opts, batchSize: batchSize, values: make(map[string]ort.Value), - withCache: withCache, - withLogits: withLogits, - withLogProbs: withLogProbs, } } func (a *Allocator) InputNames() []string { capacity := 3 - if a.withCache { + if a.options.WithCache { capacity += 2 * a.config.NumLayers } @@ -41,7 +37,7 @@ func (a *Allocator) InputNames() []string { names = append(names, "input_ids", "position_ids", "attention_mask") - if a.withCache { + if a.options.WithCache { for i := range a.config.NumLayers { names = append(names, fmt.Sprintf("past_key_values.%d.key", i), fmt.Sprintf("past_key_values.%d.value", i)) } @@ -53,29 +49,29 @@ func (a *Allocator) InputNames() []string { func (a *Allocator) OutputNames() []string { capacity := 0 - if a.withLogits { + if a.options.WithLogits { capacity++ } - if a.withLogProbs { + if a.options.WithLogProbs { capacity++ } - if a.withCache { + if a.options.WithCache { capacity += 2 * a.config.NumLayers } names := make([]string, 0, capacity) - if a.withLogits { + if a.options.WithLogits { names = append(names, "logits") } - if a.withLogProbs { + if a.options.WithLogProbs { names = append(names, "token_logprobs") } - if a.withCache { + if a.options.WithCache { for i := range a.config.NumLayers { names = append(names, fmt.Sprintf("present.%d.key", i), fmt.Sprintf("present.%d.value", i)) } @@ -119,7 +115,7 @@ func (a *Allocator) Init(tokens []int64) error { func (a *Allocator) initInputs(tokens []int64) error { capacity := 3 - if a.withCache { + if a.options.WithCache { capacity += 2 * a.config.NumLayers } @@ -143,7 +139,7 @@ func (a *Allocator) initInputs(tokens []int64) error { names = append(names, "attention_mask") - if !a.withCache { + if !a.options.WithCache { a.inputNames = names return nil @@ -171,21 +167,21 @@ func (a *Allocator) initInputs(tokens []int64) error { func (a *Allocator) initOutputs(tokens []int64) error { capacity := 0 - if a.withLogits { + if a.options.WithLogits { capacity++ } - if a.withLogProbs { + if a.options.WithLogProbs { capacity++ } - if a.withCache { + if a.options.WithCache { capacity += 2 * a.config.NumLayers } names := make([]string, 0, capacity) - if a.withLogits { + if a.options.WithLogits { if err := a.logits(false); err != nil { return err } @@ -193,7 +189,7 @@ func (a *Allocator) initOutputs(tokens []int64) error { names = append(names, "logits") } - if a.withLogProbs { + if a.options.WithLogProbs { if err := a.logProbs(false); err != nil { return err } @@ -201,7 +197,7 @@ func (a *Allocator) initOutputs(tokens []int64) error { names = append(names, "token_logprobs") } - if !a.withCache { + if !a.options.WithCache { a.outputNames = names return nil @@ -247,13 +243,13 @@ func (a *Allocator) Step(token int64) error { return err } - if a.withLogits { + if a.options.WithLogits { if err := a.logits(true); err != nil { return err } } - if a.withLogProbs { + if a.options.WithLogProbs { if err := a.logProbs(true); err != nil { return err } diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go index be8d011..a1b3cf8 100644 --- a/gpt2/cmd/gpt2/main.go +++ b/gpt2/cmd/gpt2/main.go @@ -31,7 +31,13 @@ func generate(prompt []int64) { cfg.VocabSize = 8193 - m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_cache.onnx", "0", cfg, true, true, false) + opts := gpt2.Options{ + WithCache: true, + WithLogits: true, + WithLogProbs: false, + } + + m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_cache.onnx", "0", cfg, opts) if err := m.Init(); err != nil { log.Fatal(err) @@ -55,7 +61,13 @@ func score(prompt []int64) { cfg.VocabSize = 8193 - m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_eval.onnx", "0", cfg, false, false, true) + opts := gpt2.Options{ + WithCache: false, + WithLogits: false, + WithLogProbs: true, + } + + m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_eval.onnx", "0", cfg, opts) if err := m.Init(); err != nil { log.Fatal(err) diff --git a/gpt2/config.go b/gpt2/config.go index b77578d..cd3984d 100644 --- a/gpt2/config.go +++ b/gpt2/config.go @@ -8,6 +8,12 @@ type Config struct { NumPositions int } +type Options struct { + WithCache bool + WithLogits bool + WithLogProbs bool +} + func DefaultConfig() Config { return Config{ VocabSize: 50257, 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") } diff --git a/gpt2/model_test.go b/gpt2/model_test.go index a4de628..e562e0d 100644 --- a/gpt2/model_test.go +++ b/gpt2/model_test.go @@ -40,7 +40,13 @@ func fromModel() []float32 { } func model() *Model { - m := NewModel("models/base/model.onnx", "0", DefaultConfig(), true, true, false) // TODO check if CUDA is available + opts := Options{ + WithCache: true, + WithLogits: true, + WithLogProbs: false, + } + + m := NewModel("models/base/model.onnx", "0", DefaultConfig(), opts) // TODO check if CUDA is available if err := m.Init(); err != nil { log.Fatal(err) diff --git a/llm/cmd/eval/main.go b/llm/cmd/eval/main.go index 10c7db2..88be500 100644 --- a/llm/cmd/eval/main.go +++ b/llm/cmd/eval/main.go @@ -34,7 +34,13 @@ func data() *dataset.ParquetReader { } func model() *gpt2.Model { - m := gpt2.NewModel("gpt2/models/base/model_eval.onnx", "0", gpt2.DefaultConfig(), false, false, true) + opts := gpt2.Options{ + WithCache: false, + WithLogits: false, + WithLogProbs: true, + } + + m := gpt2.NewModel("gpt2/models/base/model_eval.onnx", "0", gpt2.DefaultConfig(), opts) if err := m.Init(); err != nil { log.Fatal(err) |
