summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-06 22:37:08 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-06 22:37:08 +0200
commit42ee72bf62a03cffc0d10bb691a1b011dfadccfa (patch)
tree6fff22a5eada12fab51e9fd5846f347e486cedcb
parent10657ac011b7c32f9d272d7bd8d279edd1caf550 (diff)
Access outputs by name
-rw-r--r--gpt2/allocator.go10
-rw-r--r--gpt2/model.go12
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...)
}