diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-06 22:39:07 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-06 23:27:30 +0200 |
| commit | a8f1b564dba5cc4029ff25a49d2aa47bc7dab406 (patch) | |
| tree | e601fd866015175c740699aa1033a6e8d708eec8 | |
| parent | 42ee72bf62a03cffc0d10bb691a1b011dfadccfa (diff) | |
Add parameter to control logits output binding
| -rw-r--r-- | gpt2/allocator.go | 54 | ||||
| -rw-r--r-- | gpt2/cmd/gpt2/main.go | 4 | ||||
| -rw-r--r-- | gpt2/model.go | 12 | ||||
| -rw-r--r-- | gpt2/model_test.go | 2 | ||||
| -rw-r--r-- | llm/cmd/eval/main.go | 2 |
5 files changed, 50 insertions, 24 deletions
diff --git a/gpt2/allocator.go b/gpt2/allocator.go index a3c7c26..3736b36 100644 --- a/gpt2/allocator.go +++ b/gpt2/allocator.go @@ -13,14 +13,16 @@ type Allocator struct { outputNames []string values map[string]ort.Value withCache bool + withLogits bool withLogProbs bool } -func NewAllocator(config Config, withCache bool, withLogProbs bool) *Allocator { +func NewAllocator(config Config, withCache bool, withLogits bool, withLogProbs bool) *Allocator { return &Allocator{ config: config, values: make(map[string]ort.Value), withCache: withCache, + withLogits: withLogits, withLogProbs: withLogProbs, } } @@ -46,7 +48,15 @@ func (a *Allocator) InputNames() []string { } func (a *Allocator) OutputNames() []string { - capacity := 1 + capacity := 0 + + if a.withLogits { + capacity++ + } + + if a.withLogProbs { + capacity++ + } if a.withCache { capacity += 2 * a.config.nLayers @@ -54,10 +64,12 @@ func (a *Allocator) OutputNames() []string { names := make([]string, 0, capacity) + if a.withLogits { + names = append(names, "logits") + } + if a.withLogProbs { names = append(names, "log_probs") - } else { - names = append(names, "logits") } if a.withCache { @@ -145,7 +157,15 @@ func (a *Allocator) initInputs(tokens []int64) error { } func (a *Allocator) initOutputs(tokens []int64) error { - capacity := 1 + capacity := 0 + + if a.withLogits { + capacity++ + } + + if a.withLogProbs { + capacity++ + } if a.withCache { capacity += 2 * a.config.nLayers @@ -153,18 +173,20 @@ func (a *Allocator) initOutputs(tokens []int64) error { names := make([]string, 0, capacity) - if a.withLogProbs { - if err := a.logProbs(tokens, false); err != nil { + if a.withLogits { + if err := a.logits(tokens, false); err != nil { return err } - names = append(names, "log_probs") - } else { - if err := a.logits(tokens, false); err != nil { + names = append(names, "logits") + } + + if a.withLogProbs { + if err := a.logProbs(tokens, false); err != nil { return err } - names = append(names, "logits") + names = append(names, "log_probs") } if !a.withCache { @@ -207,12 +229,14 @@ func (a *Allocator) Step(token int64) error { return err } - if a.withLogProbs { - if err := a.logProbs(tokens, true); err != nil { + if a.withLogits { + if err := a.logits(tokens, true); err != nil { return err } - } else { - if err := a.logits(tokens, true); err != nil { + } + + if a.withLogProbs { + if err := a.logProbs(tokens, true); err != nil { return err } } diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go index 6c7a46b..663bdd8 100644 --- a/gpt2/cmd/gpt2/main.go +++ b/gpt2/cmd/gpt2/main.go @@ -19,7 +19,7 @@ func main() { } func generate(prompt []int64) { - m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_cache.onnx", "0", gpt2.NewDefaultConfig().WithVocabSize(8193), true, false) + m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_cache.onnx", "0", gpt2.NewDefaultConfig().WithVocabSize(8193), true, true, false) if err := m.Init(); err != nil { log.Fatal(err) @@ -41,7 +41,7 @@ func generate(prompt []int64) { } func score(prompt []int64) { - m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_eval.onnx", "0", gpt2.NewDefaultConfig().WithVocabSize(8193), false, true) + m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_eval.onnx", "0", gpt2.NewDefaultConfig().WithVocabSize(8193), false, false, true) if err := m.Init(); err != nil { log.Fatal(err) 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 { diff --git a/gpt2/model_test.go b/gpt2/model_test.go index 6677f73..d197ccc 100644 --- a/gpt2/model_test.go +++ b/gpt2/model_test.go @@ -26,7 +26,7 @@ func fromModel() []float32 { } func model() *Model { - m := NewModel("models/base/model.onnx", "0", NewDefaultConfig(), true, false) // TODO check if CUDA is available + m := NewModel("models/base/model.onnx", "0", NewDefaultConfig(), true, true, false) // 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 44dfedd..78a865e 100644 --- a/llm/cmd/eval/main.go +++ b/llm/cmd/eval/main.go @@ -26,7 +26,7 @@ func data() *dataset.ParquetReader { } func model() *gpt2.Model { - m := gpt2.NewModel("gpt2/models/base/model.onnx", "0", gpt2.NewDefaultConfig(), true, false) + m := gpt2.NewModel("gpt2/models/base/model.onnx", "0", gpt2.NewDefaultConfig(), true, true, false) if err := m.Init(); err != nil { log.Fatal(err) |
