summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-07 11:13:02 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-07 11:13:02 +0200
commit58ef69d71be7cdaa005a4ee21ec5ba43511e7196 (patch)
treecf370ca7c0b134837685a238a8d9ab5125e4e827
parentb81fa93e098ca689cb6f73627968fe0b4d49decb (diff)
Rename log probs output
-rw-r--r--gpt2/allocator.go8
-rw-r--r--gpt2/model.go4
2 files changed, 6 insertions, 6 deletions
diff --git a/gpt2/allocator.go b/gpt2/allocator.go
index 3736b36..2609aeb 100644
--- a/gpt2/allocator.go
+++ b/gpt2/allocator.go
@@ -69,7 +69,7 @@ func (a *Allocator) OutputNames() []string {
}
if a.withLogProbs {
- names = append(names, "log_probs")
+ names = append(names, "token_logprobs")
}
if a.withCache {
@@ -186,7 +186,7 @@ func (a *Allocator) initOutputs(tokens []int64) error {
return err
}
- names = append(names, "log_probs")
+ names = append(names, "token_logprobs")
}
if !a.withCache {
@@ -419,11 +419,11 @@ func (a *Allocator) logits(tokens []int64, force bool) error {
}
func (a *Allocator) logProbs(tokens []int64, force bool) error {
- const name = "log_probs"
+ const name = "token_logprobs"
if _, ok := a.values[name]; ok {
if !force {
- panic("log_probs already allocated")
+ panic("token_logprobs already allocated")
}
_ = a.values[name].Destroy()
diff --git a/gpt2/model.go b/gpt2/model.go
index d0d1c2f..2b23cdf 100644
--- a/gpt2/model.go
+++ b/gpt2/model.go
@@ -155,7 +155,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 requires log_probs output")
+ panic("score requires token_logprobs output")
}
if err := m.allocator.Init(tokens); err != nil {
@@ -167,7 +167,7 @@ func (m *Model) Score(tokens []int64, logProbs *[]float32) error {
}
if logProbs != nil {
- d := m.allocator.Value("log_probs").(*ort.Tensor[float32]).GetData()
+ d := m.allocator.Value("token_logprobs").(*ort.Tensor[float32]).GetData()
*logProbs = append(*logProbs, d...)
}