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 /gpt2/model.go | |
| parent | b81fa93e098ca689cb6f73627968fe0b4d49decb (diff) | |
Rename log probs output
Diffstat (limited to 'gpt2/model.go')
| -rw-r--r-- | gpt2/model.go | 4 |
1 files changed, 2 insertions, 2 deletions
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...) } |
