summaryrefslogtreecommitdiff
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
parent42ee72bf62a03cffc0d10bb691a1b011dfadccfa (diff)
Add parameter to control logits output binding
-rw-r--r--gpt2/allocator.go54
-rw-r--r--gpt2/cmd/gpt2/main.go4
-rw-r--r--gpt2/model.go12
-rw-r--r--gpt2/model_test.go2
-rw-r--r--llm/cmd/eval/main.go2
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)