summaryrefslogtreecommitdiff
path: root/gpt2
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2025-11-24 23:41:41 +0100
committerJonas Knobloch <jonas.knobloch@t-online.de>2025-11-26 18:23:05 +0100
commita29505606ecdd0db353d7d2c1c0a79f24e0ac9d7 (patch)
tree9a68c13dc9825f253757507f0e8ce2286ea32dce /gpt2
parent9077d8f56f245ab1838cb9b9671e1cd2c6b568fa (diff)
Use seen tokens to warm up cache
Diffstat (limited to 'gpt2')
-rw-r--r--gpt2/model.go124
1 files changed, 77 insertions, 47 deletions
diff --git a/gpt2/model.go b/gpt2/model.go
index 1058ef4..2dcb452 100644
--- a/gpt2/model.go
+++ b/gpt2/model.go
@@ -95,57 +95,77 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in
return nil, errors.New("empty prompt")
}
- context := int64(len(prompt))
+ n := int64(len(prompt))
+ r := make([]int64, steps)
- cacheValues := emptyCache()
+ outputs := initCache()
- defer func() {
- destroyValues(cacheValues)
- }()
-
- token := prompt[0]
-
- out := make([]int64, 0, steps+1)
+ if o, err := m.forward(prompt, 0, outputs); err != nil {
+ defer destroyValues(outputs)
- for step := range context + steps {
- outputs, err := m.forward(token, step, cacheValues)
+ return nil, err
+ } else {
+ destroyValues(outputs)
- if err != nil {
- return nil, err
- }
+ outputs = o
+ }
- l := outputs[0].(*ort.Tensor[float32]).GetData()
+ for step := range steps {
+ l := m.logits(outputs[0])
if logits != nil {
- *logits = append(*logits, l)
+ *logits = append(*logits, l...)
}
- idx, _ := topK(softmax(l), 5)
+ next := m.sample(l[len(l)-1])
- // fmt.Printf("\n%d\n\n", token)
+ r[step] = next
+
+ _ = outputs[0].Destroy()
- // for i, t := range idx {
- // fmt.Printf("%.4f %.4f [%d]\n", l[t], p[i], t)
- // }
+ if o, err := m.forward([]int64{next}, n+step, outputs[1:]); err != nil {
+ defer destroyValues(outputs[1:])
- if step < context-1 {
- token = prompt[step+1]
+ return nil, err
} else {
- token = int64(idx[0]) // choose best token
- out = append(out, token)
+ destroyValues(outputs[1:])
+
+ outputs = o
+ }
+ }
+
+ if logits != nil {
+ for _, l := range m.logits(outputs[0]) {
+ *logits = append(*logits, l)
}
+ }
- _ = outputs[0].Destroy()
+ _ = outputs[0].Destroy()
+
+ return r, nil
+}
+
+func (m *Model) logits(output ort.Value) [][]float32 {
+ d := output.(*ort.Tensor[float32]).GetData()
+ n := len(d) / vocabSize
+ l := make([][]float32, n)
- destroyValues(cacheValues)
+ for i := range n {
+ s := i * vocabSize
- cacheValues = outputs[1:]
+ l[i] = d[s : s+vocabSize : s+vocabSize]
}
- return out[:steps], nil
+ return l
}
-func (m *Model) forward(token int64, position int64, cacheValues []ort.Value) ([]ort.Value, error) {
+func (m *Model) sample(logits []float32) int64 {
+ idx, _ := topK(softmax(logits), 5)
+
+ return int64(idx[0])
+}
+
+func (m *Model) forward(tokens []int64, start int64, cache []ort.Value) ([]ort.Value, error) {
var binding *ort.IoBinding
if b, err := m.session.CreateIoBinding(); err != nil {
@@ -159,7 +179,7 @@ func (m *Model) forward(token int64, position int64, cacheValues []ort.Value) ([
var inputs []ort.Value
var outputs []ort.Value
- if in, err := initInputs(token, position); err != nil {
+ if in, err := initInputs(tokens, start); err != nil {
return nil, err
} else {
inputs = in
@@ -167,13 +187,13 @@ func (m *Model) forward(token int64, position int64, cacheValues []ort.Value) ([
defer destroyValues(inputs)
- if out, err := initOutputs(position); err != nil {
+ if out, err := initOutputs(tokens, start); err != nil {
return nil, err
} else {
outputs = out
}
- inputs = append(inputs, cacheValues...)
+ inputs = append(inputs, cache...)
if len(inputs) != len(m.inputNames) {
panic("unexpected input length")
@@ -218,7 +238,7 @@ func destroyValues(values []ort.Value) {
}
}
-func emptyCache() []ort.Value {
+func initCache() []ort.Value {
values := make([]ort.Value, 0, 2*nLayers)
shape := []int64{1, int64(nHeads), 0, int64(headDim)}
@@ -232,25 +252,35 @@ func emptyCache() []ort.Value {
return values
}
-func initInputs(token, position int64) ([]ort.Value, error) {
- var tokens *ort.Tensor[int64]
- var positions *ort.Tensor[int64]
+func initInputs(tokens []int64, start int64) ([]ort.Value, error) {
+ var inputIDs *ort.Tensor[int64]
+ var positionIDs *ort.Tensor[int64]
var attentionMask *ort.Tensor[int64]
- if t, err := ort.NewTensor[int64]([]int64{1, 1}, []int64{token}); err != nil {
+ inputsShape := []int64{1, int64(len(tokens))}
+ inputsData := tokens
+
+ if t, err := ort.NewTensor[int64](inputsShape, inputsData); err != nil {
return nil, err
} else {
- tokens = t
+ 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]([]int64{1, 1}, []int64{position}); err != nil {
+ if p, err := ort.NewTensor[int64](positionsShape, positionsData); err != nil {
return nil, err
} else {
- positions = p
+ positionIDs = p
}
- maskData := make([]int64, position+1)
- maskShape := []int64{1, position + 1}
+ maskData := make([]int64, start+int64(len(tokens)))
+ maskShape := []int64{1, start + int64(len(tokens))}
for i := range maskData {
maskData[i] = 1
@@ -262,15 +292,15 @@ func initInputs(token, position int64) ([]ort.Value, error) {
attentionMask = m
}
- inputs := []ort.Value{ort.Value(tokens), ort.Value(positions), ort.Value(attentionMask)}
+ inputs := []ort.Value{ort.Value(inputIDs), ort.Value(positionIDs), ort.Value(attentionMask)}
return inputs, nil
}
-func initOutputs(position int64) ([]ort.Value, error) {
+func initOutputs(tokens []int64, start int64) ([]ort.Value, error) {
outputs := make([]ort.Value, 0, 1+2*nLayers)
- logits, err := ort.NewEmptyTensor[float32]([]int64{1, 1, int64(vocabSize)})
+ logits, err := ort.NewEmptyTensor[float32]([]int64{1, int64(len(tokens)), int64(vocabSize)})
if err != nil {
return nil, err
@@ -278,7 +308,7 @@ func initOutputs(position int64) ([]ort.Value, error) {
outputs = append(outputs, ort.Value(logits))
- shape := []int64{1, int64(nHeads), position + 1, int64(headDim)}
+ shape := []int64{1, int64(nHeads), start + int64(len(tokens)), int64(headDim)}
for range nLayers {
kTensor, _ := ort.NewEmptyTensor[float32](shape)