diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-04 21:22:59 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-04 21:22:59 +0200 |
| commit | 758e27cf2b3714588af7db1f5da057b15995ba26 (patch) | |
| tree | 9f7286ef2bf0d458fec8d04a0ff81811b71e86a9 | |
| parent | 7696b62cf3d3d7ea482b7028a49d06522262da4d (diff) | |
Add allocator abstraction
| -rw-r--r-- | gpt2/allocator.go | 400 | ||||
| -rw-r--r-- | gpt2/model.go | 209 |
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 { |
