diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-07 11:13:02 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-07 11:13:02 +0200 |
| commit | 58ef69d71be7cdaa005a4ee21ec5ba43511e7196 (patch) | |
| tree | cf370ca7c0b134837685a238a8d9ab5125e4e827 | |
| parent | b81fa93e098ca689cb6f73627968fe0b4d49decb (diff) | |
Rename log probs output
| -rw-r--r-- | gpt2/allocator.go | 8 | ||||
| -rw-r--r-- | 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...) } |
