From 42ee72bf62a03cffc0d10bb691a1b011dfadccfa Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Mon, 6 Apr 2026 22:37:08 +0200 Subject: Access outputs by name --- gpt2/allocator.go | 10 ++++++++++ gpt2/model.go | 12 +++--------- 2 files changed, 13 insertions(+), 9 deletions(-) diff --git a/gpt2/allocator.go b/gpt2/allocator.go index 30fb621..a3c7c26 100644 --- a/gpt2/allocator.go +++ b/gpt2/allocator.go @@ -256,6 +256,16 @@ func (a *Allocator) Outputs() ([]string, []ort.Value) { return a.outputNames, vals } +func (a *Allocator) Value(name string) ort.Value { + v, ok := a.values[name] + + if !ok { + panic("unknown value: " + name) + } + + return v +} + func (a *Allocator) inputIDs(tokens []int64, force bool) error { const name = "input_ids" diff --git a/gpt2/model.go b/gpt2/model.go index 87a7e09..214affe 100644 --- a/gpt2/model.go +++ b/gpt2/model.go @@ -127,9 +127,7 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in r := make([]int64, steps) for step := range steps { - _, outputs := m.allocator.Outputs() - - l := m.logits(outputs[0]) + l := m.logits(m.allocator.Value("logits")) if logits != nil { *logits = append(*logits, l...) @@ -149,9 +147,7 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in } if logits != nil { - _, outVals := m.allocator.Outputs() - - for _, l := range m.logits(outVals[0]) { + for _, l := range m.logits(m.allocator.Value("logits")) { *logits = append(*logits, l) } } @@ -173,9 +169,7 @@ func (m *Model) Score(tokens []int64, logProbs *[]float32) error { } if logProbs != nil { - _, outputs := m.allocator.Outputs() - - d := outputs[0].(*ort.Tensor[float32]).GetData() + d := m.allocator.Value("log_probs").(*ort.Tensor[float32]).GetData() *logProbs = append(*logProbs, d...) } -- cgit v1.2.3