summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-04 21:22:59 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-04 21:22:59 +0200
commit758e27cf2b3714588af7db1f5da057b15995ba26 (patch)
tree9f7286ef2bf0d458fec8d04a0ff81811b71e86a9
parent7696b62cf3d3d7ea482b7028a49d06522262da4d (diff)
Add allocator abstraction
-rw-r--r--gpt2/allocator.go400
-rw-r--r--gpt2/model.go209
2 files changed, 432 insertions, 177 deletions
diff --git a/gpt2/allocator.go b/gpt2/allocator.go
new file mode 100644
index 0000000..e65c752
--- /dev/null
+++ b/gpt2/allocator.go
@@ -0,0 +1,400 @@
+package gpt2
+
+import (
+ "fmt"
+
+ ort "github.com/yalue/onnxruntime_go"
+)
+
+type Allocator struct {
+ step int64
+ inputNames []string
+ outputNames []string
+ values map[string]ort.Value
+ withCache bool
+}
+
+func NewAllocator(withCache bool) *Allocator {
+ return &Allocator{
+ values: make(map[string]ort.Value),
+ withCache: withCache,
+ }
+}
+
+func (a *Allocator) InputNames() []string {
+ names := make([]string, 0, 3+2*nLayers)
+
+ names = append(names, "input_ids", "position_ids", "attention_mask")
+
+ if a.withCache {
+ for i := range nLayers {
+ names = append(names, fmt.Sprintf("past_key_values.%d.key", i), fmt.Sprintf("past_key_values.%d.value", i))
+ }
+ }
+
+ return names
+}
+
+func (a *Allocator) OutputNames() []string {
+ names := make([]string, 0, 1+2*nLayers)
+
+ names = append(names, "logits")
+
+ if a.withCache {
+ for i := range nLayers {
+ names = append(names, fmt.Sprintf("present.%d.key", i), fmt.Sprintf("present.%d.value", i))
+ }
+ }
+
+ return names
+}
+
+func (a *Allocator) Init(tokens []int64) error {
+ a.Destroy()
+
+ a.values = make(map[string]ort.Value)
+ a.step = 0
+
+ if err := a.initInputs(tokens); err != nil {
+ a.Destroy()
+
+ return err
+ }
+
+ if err := a.initOutputs(tokens); err != nil {
+ a.Destroy()
+
+ return err
+ }
+
+ a.step = int64(len(tokens))
+
+ return nil
+}
+
+func (a *Allocator) initInputs(tokens []int64) error {
+ capacity := 3
+
+ if a.withCache {
+ capacity += 2 * nLayers
+ }
+
+ names := make([]string, 0, capacity)
+
+ if err := a.inputIDs(tokens, false); err != nil {
+ return err
+ }
+
+ names = append(names, "input_ids")
+
+ if err := a.positionIDs(tokens, 0, false); err != nil {
+ return err
+ }
+
+ names = append(names, "position_ids")
+
+ if err := a.attentionMask(tokens, 0, false); err != nil {
+ return err
+ }
+
+ names = append(names, "attention_mask")
+
+ if !a.withCache {
+ a.inputNames = names
+
+ return nil
+ }
+
+ for i := range int64(nLayers) {
+ if err := a.pastKeyValues(i, "key", false); err != nil {
+ return err
+ }
+
+ names = append(names, fmt.Sprintf("past_key_values.%d.key", i))
+
+ if err := a.pastKeyValues(i, "value", false); err != nil {
+ return err
+ }
+
+ names = append(names, fmt.Sprintf("past_key_values.%d.value", i))
+ }
+
+ a.inputNames = names
+
+ return nil
+}
+
+func (a *Allocator) initOutputs(tokens []int64) error {
+ capacity := 1
+
+ if a.withCache {
+ capacity += 2 * nLayers
+ }
+
+ names := make([]string, 0, capacity)
+
+ if err := a.logits(tokens, false); err != nil {
+ return err
+ }
+
+ names = append(names, "logits")
+
+ if !a.withCache {
+ a.outputNames = names
+
+ return nil
+ }
+
+ for i := range int64(nLayers) {
+ if err := a.presentKeyValues(tokens, 0, i, "key", false); err != nil {
+ return err
+ }
+
+ names = append(names, fmt.Sprintf("present.%d.key", i))
+
+ if err := a.presentKeyValues(tokens, 0, i, "value", false); err != nil {
+ return err
+ }
+
+ names = append(names, fmt.Sprintf("present.%d.value", i))
+ }
+
+ a.outputNames = names
+
+ return nil
+}
+
+func (a *Allocator) Step(token int64) error {
+ tokens := []int64{token}
+
+ if err := a.inputIDs(tokens, true); err != nil {
+ return err
+ }
+
+ if err := a.positionIDs(tokens, a.step, true); err != nil {
+ return err
+ }
+
+ if err := a.attentionMask(tokens, a.step, true); err != nil {
+ return err
+ }
+
+ if err := a.logits(tokens, true); err != nil {
+ return err
+ }
+
+ for i := range int64(nLayers) {
+ for _, suffix := range []string{"key", "value"} {
+ if err := a.rotateCache(tokens, i, suffix); err != nil {
+ return err
+ }
+ }
+ }
+
+ a.step++
+
+ return nil
+}
+
+func (a *Allocator) Destroy() {
+ for _, v := range a.values {
+ _ = v.Destroy()
+ }
+}
+
+func (a *Allocator) Inputs() ([]string, []ort.Value) {
+ vals := make([]ort.Value, len(a.inputNames))
+
+ for i, n := range a.inputNames {
+ vals[i] = a.values[n]
+ }
+
+ return a.inputNames, vals
+}
+
+func (a *Allocator) Outputs() ([]string, []ort.Value) {
+ vals := make([]ort.Value, len(a.outputNames))
+
+ for i, n := range a.outputNames {
+ vals[i] = a.values[n]
+ }
+
+ return a.outputNames, vals
+}
+
+func (a *Allocator) inputIDs(tokens []int64, force bool) error {
+ const name = "input_ids"
+
+ if _, ok := a.values[name]; ok {
+ if !force {
+ panic("input_ids already allocated")
+ }
+
+ _ = a.values[name].Destroy()
+ }
+
+ shape := []int64{1, int64(len(tokens))}
+
+ if t, err := ort.NewTensor[int64](shape, tokens); err != nil {
+ return err
+ } else {
+ a.values[name] = ort.Value(t)
+ }
+
+ return nil
+}
+
+func (a *Allocator) positionIDs(tokens []int64, start int64, force bool) error {
+ const name = "position_ids"
+
+ if _, ok := a.values[name]; ok {
+ if !force {
+ panic("position_ids already allocated")
+ }
+
+ _ = a.values[name].Destroy()
+ }
+
+ data := make([]int64, len(tokens))
+ shape := []int64{1, int64(len(tokens))}
+
+ for i := range len(tokens) {
+ data[i] = start + int64(i)
+ }
+
+ if t, err := ort.NewTensor[int64](shape, data); err != nil {
+ return err
+ } else {
+ a.values[name] = ort.Value(t)
+ }
+
+ return nil
+}
+
+func (a *Allocator) attentionMask(tokens []int64, start int64, force bool) error {
+ const name = "attention_mask"
+
+ if _, ok := a.values[name]; ok {
+ if !force {
+ panic("attention_mask already allocated")
+ }
+
+ _ = a.values[name].Destroy()
+ }
+
+ data := make([]int64, start+int64(len(tokens)))
+ shape := []int64{1, start + int64(len(tokens))}
+
+ for i := range data {
+ data[i] = 1
+ }
+
+ if t, err := ort.NewTensor[int64](shape, data); err != nil {
+ return err
+ } else {
+ a.values[name] = ort.Value(t)
+ }
+
+ return nil
+}
+
+func (a *Allocator) pastKeyValues(i int64, suffix string, force bool) error {
+ if i > nLayers {
+ panic("invalid layer index")
+ }
+
+ if suffix != "key" && suffix != "value" {
+ panic("invalid suffix")
+ }
+
+ name := fmt.Sprintf("past_key_values.%d.%s", i, suffix)
+
+ if _, ok := a.values[name]; ok {
+ if !force {
+ panic(name + " already allocated")
+ }
+
+ _ = a.values[name].Destroy()
+ }
+
+ shape := []int64{1, int64(nHeads), 0, int64(headDim)}
+
+ if t, err := ort.NewEmptyTensor[float32](shape); err != nil {
+ return err
+ } else {
+ a.values[name] = ort.Value(t)
+ }
+
+ return nil
+}
+
+func (a *Allocator) logits(tokens []int64, force bool) error {
+ const name = "logits"
+
+ if _, ok := a.values[name]; ok {
+ if !force {
+ panic("logits already allocated")
+ }
+
+ _ = a.values[name].Destroy()
+ }
+
+ shape := []int64{1, int64(len(tokens)), int64(vocabSize)}
+
+ if t, err := ort.NewEmptyTensor[float32](shape); err != nil {
+ return err
+ } else {
+ a.values[name] = ort.Value(t)
+ }
+
+ return nil
+}
+
+func (a *Allocator) presentKeyValues(tokens []int64, start, i int64, suffix string, force bool) error {
+ if i > nLayers {
+ panic("invalid layer index")
+ }
+
+ if suffix != "key" && suffix != "value" {
+ panic("invalid suffix")
+ }
+
+ name := fmt.Sprintf("present.%d.%s", i, suffix)
+
+ if _, ok := a.values[name]; ok {
+ if !force {
+ panic(name + " already allocated")
+ }
+
+ _ = a.values[name].Destroy()
+ }
+
+ shape := []int64{1, int64(nHeads), start + int64(len(tokens)), int64(headDim)}
+
+ if t, err := ort.NewEmptyTensor[float32](shape); err != nil {
+ return err
+ } else {
+ a.values[name] = ort.Value(t)
+ }
+
+ return nil
+}
+
+func (a *Allocator) rotateCache(tokens []int64, i int64, suffix string) error {
+ past := fmt.Sprintf("past_key_values.%d.%s", i, suffix)
+ present := fmt.Sprintf("present.%d.%s", i, suffix)
+
+ if v, ok := a.values[past]; ok {
+ _ = v.Destroy()
+ }
+
+ a.values[past] = a.values[present]
+
+ delete(a.values, present)
+
+ if err := a.presentKeyValues(tokens, a.step, i, suffix, false); err != nil {
+ return err
+ }
+
+ return nil
+}
diff --git a/gpt2/model.go b/gpt2/model.go
index 6161ad3..51889be 100644
--- a/gpt2/model.go
+++ b/gpt2/model.go
@@ -2,7 +2,6 @@ package gpt2
import (
"errors"
- "fmt"
"math"
"os"
"slices"
@@ -21,11 +20,10 @@ const (
)
type Model struct {
- name string
- deviceID string
- session *ort.DynamicAdvancedSession
- inputNames []string
- outputNames []string
+ name string
+ deviceID string
+ session *ort.DynamicAdvancedSession
+ allocator *Allocator
}
func NewModel(name, deviceID string) *Model {
@@ -68,19 +66,7 @@ func (m *Model) Init() error {
return err
}
- inputNames := make([]string, 0, 3+2*nLayers)
- outputNames := make([]string, 0, 1+2*nLayers)
-
- inputNames = append(inputNames, "input_ids", "position_ids", "attention_mask")
- outputNames = append(outputNames, "logits")
-
- for i := range nLayers {
- inputNames = append(inputNames, fmt.Sprintf("past_key_values.%d.key", i), fmt.Sprintf("past_key_values.%d.value", i))
- outputNames = append(outputNames, fmt.Sprintf("present.%d.key", i), fmt.Sprintf("present.%d.value", i))
- }
-
- m.inputNames = inputNames
- m.outputNames = outputNames
+ m.allocator = NewAllocator(true)
var options *ort.SessionOptions
@@ -104,7 +90,7 @@ func (m *Model) Init() error {
}
}
- if s, err := ort.NewDynamicAdvancedSession(m.name, m.inputNames, m.outputNames, options); err != nil {
+ if s, err := ort.NewDynamicAdvancedSession(m.name, m.allocator.InputNames(), m.allocator.OutputNames(), options); err != nil {
return err
} else {
m.session = s
@@ -114,6 +100,8 @@ func (m *Model) Init() error {
}
func (m *Model) Destroy() error {
+ m.allocator.Destroy()
+
return ort.DestroyEnvironment()
}
@@ -126,22 +114,19 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in
return nil, errors.New("sequence length exceeds context limit")
}
- n := int64(len(prompt))
- r := make([]int64, steps)
-
- outputs := initCache()
-
- if o, err := m.forward(prompt, 0, outputs); err != nil {
- defer destroyValues(outputs)
-
+ if err := m.allocator.Init(prompt); err != nil {
return nil, err
- } else {
- destroyValues(outputs)
+ }
- outputs = o
+ if err := m.forward(m.allocator); err != nil {
+ return nil, err
}
+ r := make([]int64, steps)
+
for step := range steps {
+ _, outputs := m.allocator.Outputs()
+
l := m.logits(outputs[0])
if logits != nil {
@@ -152,27 +137,23 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in
r[step] = next
- _ = outputs[0].Destroy()
-
- if o, err := m.forward([]int64{next}, n+step, outputs[1:]); err != nil {
- defer destroyValues(outputs[1:])
-
+ if err := m.allocator.Step(next); err != nil {
return nil, err
- } else {
- destroyValues(outputs[1:])
+ }
- outputs = o
+ if err := m.forward(m.allocator); err != nil {
+ return nil, err
}
}
if logits != nil {
- for _, l := range m.logits(outputs[0]) {
+ _, outVals := m.allocator.Outputs()
+
+ for _, l := range m.logits(outVals[0]) {
*logits = append(*logits, l)
}
}
- _ = outputs[0].Destroy()
-
return r, nil
}
@@ -196,159 +177,33 @@ func (m *Model) sample(logits []float32) int64 {
return int64(idx[0])
}
-func (m *Model) forward(tokens []int64, start int64, cache []ort.Value) ([]ort.Value, error) {
+func (m *Model) forward(alloc *Allocator) error {
var binding *ort.IoBinding
if b, err := m.session.CreateIoBinding(); err != nil {
- return nil, err
+ return err
} else {
binding = b
defer binding.Destroy()
}
- var inputs []ort.Value
- var outputs []ort.Value
+ inputNames, inputs := alloc.Inputs()
+ outputNames, outputs := alloc.Outputs()
- if in, err := initInputs(tokens, start); err != nil {
- return nil, err
- } else {
- inputs = in
- }
-
- defer destroyValues(inputs)
-
- if out, err := initOutputs(tokens, start); err != nil {
- return nil, err
- } else {
- outputs = out
- }
-
- inputs = append(inputs, cache...)
-
- if len(inputs) != len(m.inputNames) {
- panic("unexpected input length")
- }
-
- for i, name := range m.inputNames {
+ for i, name := range inputNames {
if err := binding.BindInput(name, inputs[i]); err != nil {
- return nil, err
- }
- }
-
- var ok bool
-
- defer func() {
- if !ok {
- destroyValues(outputs)
+ return err
}
- }()
-
- if len(outputs) != len(m.outputNames) {
- panic("unexpected output length")
}
- for i, name := range m.outputNames {
+ for i, name := range outputNames {
if err := binding.BindOutput(name, outputs[i]); err != nil {
- return nil, err
+ return err
}
}
- if err := m.session.RunWithBinding(binding); err != nil {
- return nil, err
- }
-
- ok = true
-
- return outputs, nil
-}
-
-func destroyValues(values []ort.Value) {
- for _, v := range values {
- _ = v.Destroy()
- }
-}
-
-func initCache() []ort.Value {
- values := make([]ort.Value, 0, 2*nLayers)
- shape := []int64{1, int64(nHeads), 0, int64(headDim)}
-
- for range nLayers {
- kTensor, _ := ort.NewEmptyTensor[float32](shape)
- vTensor, _ := ort.NewEmptyTensor[float32](shape)
-
- values = append(values, ort.Value(kTensor), ort.Value(vTensor))
- }
-
- return values
-}
-
-func initInputs(tokens []int64, start int64) ([]ort.Value, error) {
- var inputIDs *ort.Tensor[int64]
- var positionIDs *ort.Tensor[int64]
- var attentionMask *ort.Tensor[int64]
-
- inputsShape := []int64{1, int64(len(tokens))}
- inputsData := tokens
-
- if t, err := ort.NewTensor[int64](inputsShape, inputsData); err != nil {
- return nil, err
- } else {
- inputIDs = t
- }
-
- positionsShape := []int64{1, int64(len(tokens))}
- positionsData := make([]int64, len(tokens))
-
- for i := range len(tokens) {
- positionsData[i] = start + int64(i)
- }
-
- if p, err := ort.NewTensor[int64](positionsShape, positionsData); err != nil {
- return nil, err
- } else {
- positionIDs = p
- }
-
- maskData := make([]int64, start+int64(len(tokens)))
- maskShape := []int64{1, start + int64(len(tokens))}
-
- for i := range maskData {
- maskData[i] = 1
- }
-
- if m, err := ort.NewTensor[int64](maskShape, maskData); err != nil {
- return nil, err
- } else {
- attentionMask = m
- }
-
- inputs := []ort.Value{ort.Value(inputIDs), ort.Value(positionIDs), ort.Value(attentionMask)}
-
- return inputs, nil
-}
-
-func initOutputs(tokens []int64, start int64) ([]ort.Value, error) {
- outputs := make([]ort.Value, 0, 1+2*nLayers)
-
- logits, err := ort.NewEmptyTensor[float32]([]int64{1, int64(len(tokens)), int64(vocabSize)})
-
- if err != nil {
- return nil, err
- }
-
- outputs = append(outputs, ort.Value(logits))
-
- shape := []int64{1, int64(nHeads), start + int64(len(tokens)), int64(headDim)}
-
- for range nLayers {
- kTensor, _ := ort.NewEmptyTensor[float32](shape)
- vTensor, _ := ort.NewEmptyTensor[float32](shape)
-
- outputs = append(outputs, ort.Value(kTensor), ort.Value(vTensor))
- }
-
- return outputs, nil
+ return m.session.RunWithBinding(binding)
}
func softmax(logits []float32) []float32 {