diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-06 22:37:08 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-06 22:37:08 +0200 |
| commit | 42ee72bf62a03cffc0d10bb691a1b011dfadccfa (patch) | |
| tree | 6fff22a5eada12fab51e9fd5846f347e486cedcb | |
| parent | 10657ac011b7c32f9d272d7bd8d279edd1caf550 (diff) | |
Access outputs by name
| -rw-r--r-- | gpt2/allocator.go | 10 | ||||
| -rw-r--r-- | 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...) } |
