From 58ef69d71be7cdaa005a4ee21ec5ba43511e7196 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Tue, 7 Apr 2026 11:13:02 +0200 Subject: Rename log probs output --- gpt2/allocator.go | 8 ++++---- gpt2/model.go | 4 ++-- 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...) } -- cgit v1.2.3