summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-08 22:24:10 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-08 23:46:10 +0200
commit4386784c1006321c5cf0e1e005b19e0a7d7563ae (patch)
tree3d22360aaac8e4d64ca90e39b7247942073582d1
parenta006507a24b93a6af7e3cff8833fb96b48b159d7 (diff)
Refactor model options
* Add options struct
-rw-r--r--gpt2/allocator.go46
-rw-r--r--gpt2/cmd/gpt2/main.go16
-rw-r--r--gpt2/config.go6
-rw-r--r--gpt2/model.go18
-rw-r--r--gpt2/model_test.go8
-rw-r--r--llm/cmd/eval/main.go8
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)