summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--main.go43
-rw-r--r--scripts/conv.py2
2 files changed, 39 insertions, 6 deletions
diff --git a/main.go b/main.go
index 319a6b5..2ab2e1a 100644
--- a/main.go
+++ b/main.go
@@ -9,6 +9,12 @@ import (
ort "github.com/yalue/onnxruntime_go"
)
+const (
+ nLayers = 12
+ nHeads = 12
+ headDim = 64
+)
+
func main() {
ort.SetSharedLibraryPath("lib/onnxruntime-osx-arm64-1.22.0/lib/libonnxruntime.1.22.0.dylib")
@@ -24,14 +30,22 @@ func main() {
tokens, _ := ort.NewTensor[int64]([]int64{1, 1}, []int64{464})
positions, _ := ort.NewTensor[int64]([]int64{1, 1}, []int64{0})
attentionMask, _ := ort.NewTensor[int64]([]int64{1, 1}, []int64{1})
- output, _ := ort.NewEmptyTensor[float32]([]int64{1, 1, 50257})
+ logits, _ := ort.NewEmptyTensor[float32]([]int64{1, 1, 50257})
+
+ inputs := []ort.Value{ort.Value(tokens), ort.Value(positions), ort.Value(attentionMask)}
+ outputs := []ort.Value{ort.Value(logits)}
+
+ cacheNames, cacheValues := emptyCache()
+
+ inputNames = append(inputNames, cacheNames...)
+ inputs = append(inputs, cacheValues...)
session, err := ort.NewAdvancedSession(
"scripts/onnx-gpt2/model.onnx",
inputNames,
outputNames,
- []ort.Value{ort.Value(tokens), ort.Value(positions), ort.Value(attentionMask)},
- []ort.Value{ort.Value(output)},
+ inputs,
+ outputs,
nil,
)
@@ -45,8 +59,8 @@ func main() {
log.Fatal(err)
}
- logits := output.GetData()
- probs := softmax(logits)
+ data := logits.GetData()
+ probs := softmax(data)
idx, p := topK(probs, 10)
@@ -55,6 +69,25 @@ func main() {
}
}
+func emptyCache() ([]string, []ort.Value) {
+ names := make([]string, 0, 2*nLayers)
+ values := make([]ort.Value, 0, 2*nLayers)
+ shape := []int64{1, int64(nHeads), 0, int64(headDim)}
+
+ for i := range nLayers {
+ kName := fmt.Sprintf("past_key_values.%d.key", i)
+ vName := fmt.Sprintf("past_key_values.%d.value", i)
+
+ kTensor, _ := ort.NewEmptyTensor[float32](shape)
+ vTensor, _ := ort.NewEmptyTensor[float32](shape)
+
+ names = append(names, kName, vName)
+ values = append(values, ort.Value(kTensor), ort.Value(vTensor))
+ }
+
+ return names, values
+}
+
func softmax(logits []float32) []float32 {
m := logits[0]
diff --git a/scripts/conv.py b/scripts/conv.py
index 07d63f9..cf0cbb8 100644
--- a/scripts/conv.py
+++ b/scripts/conv.py
@@ -11,5 +11,5 @@ from transformers import AutoTokenizer, AutoModelForCausalLM
from optimum.onnxruntime import ORTModelForCausalLM
model_id = "gpt2"
-model = ORTModelForCausalLM.from_pretrained(model_id, export=True, use_cache=False)
+model = ORTModelForCausalLM.from_pretrained(model_id, export=True, use_cache=True)
model.save_pretrained("onnx-gpt2") \ No newline at end of file