summaryrefslogtreecommitdiff
path: root/gpt2
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2025-11-13 19:25:27 +0100
committerJonas Knobloch <jonas.knobloch@t-online.de>2025-11-13 20:10:37 +0100
commit019a25a6082a8cde372304e34dcd0ae3d5e875ed (patch)
tree908157de7d20b67d8436a4fbed933c59f4a5d16e /gpt2
parent148d5910f3577d7f8ed6aef57416a5f9d17efdc6 (diff)
Refactor into workspace
Diffstat (limited to 'gpt2')
-rw-r--r--gpt2/cmd/gpt2/main.go27
-rw-r--r--gpt2/go.mod5
-rw-r--r--gpt2/go.sum2
-rw-r--r--gpt2/main.go232
-rw-r--r--gpt2/scripts/conv.py15
5 files changed, 281 insertions, 0 deletions
diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go
new file mode 100644
index 0000000..ff84f42
--- /dev/null
+++ b/gpt2/cmd/gpt2/main.go
@@ -0,0 +1,27 @@
+package main
+
+import (
+ "fmt"
+ "gpt2"
+ "log"
+
+ ort "github.com/yalue/onnxruntime_go"
+)
+
+func main() {
+ ort.SetSharedLibraryPath("lib/onnxruntime-osx-arm64-1.22.0/lib/libonnxruntime.1.22.0.dylib")
+
+ if err := ort.InitializeEnvironment(); err != nil {
+ log.Fatal(err)
+ }
+
+ defer ort.DestroyEnvironment()
+
+ prompt := []int64{464, 2068, 7586, 21831}
+
+ if out, err := gpt2.Generate("scripts/onnx-gpt2/model.onnx", prompt, 5, nil); err != nil {
+ log.Fatal(err)
+ } else {
+ fmt.Printf("\n%v\n", out)
+ }
+}
diff --git a/gpt2/go.mod b/gpt2/go.mod
new file mode 100644
index 0000000..2ca372e
--- /dev/null
+++ b/gpt2/go.mod
@@ -0,0 +1,5 @@
+module gpt2
+
+go 1.24
+
+require github.com/yalue/onnxruntime_go v1.22.0
diff --git a/gpt2/go.sum b/gpt2/go.sum
new file mode 100644
index 0000000..f2c6460
--- /dev/null
+++ b/gpt2/go.sum
@@ -0,0 +1,2 @@
+github.com/yalue/onnxruntime_go v1.22.0 h1:SzqOfFRRrLRRAFR5VoSxABjTiQSAi8Y4ETYKrMFK1jk=
+github.com/yalue/onnxruntime_go v1.22.0/go.mod h1:b4X26A8pekNb1ACJ58wAXgNKeUCGEAQ9dmACut9Sm/4=
diff --git a/gpt2/main.go b/gpt2/main.go
new file mode 100644
index 0000000..ee9e19f
--- /dev/null
+++ b/gpt2/main.go
@@ -0,0 +1,232 @@
+package gpt2
+
+import (
+ "errors"
+ "fmt"
+ "log"
+ "math"
+ "sort"
+
+ ort "github.com/yalue/onnxruntime_go"
+)
+
+const (
+ vocabSize = 50257
+ nLayers = 12
+ nHeads = 12
+ headDim = 64
+)
+
+func Generate(model string, prompt []int64, steps int64, logits *[][]float32) ([]int64, error) {
+ if len(prompt) == 0 {
+ return nil, errors.New("empty prompt")
+ }
+
+ context := int64(len(prompt))
+
+ cacheNames, cacheValues := emptyCache()
+
+ token := prompt[0]
+
+ out := make([]int64, 0, steps+1)
+
+ for step := range context + steps {
+ _, _, outputs, err := forward(model, token, step, cacheNames, cacheValues)
+
+ if err != nil {
+ return nil, err
+ }
+
+ l := outputs[0].(*ort.Tensor[float32]).GetData()
+
+ if logits != nil {
+ *logits = append(*logits, l)
+ }
+
+ idx, p := topK(softmax(l), 5)
+
+ fmt.Printf("\n%d\n\n", token)
+
+ for i, t := range idx {
+ fmt.Printf("%.4f %.4f [%d]\n", l[t], p[i], t)
+ }
+
+ if step < context-1 {
+ token = prompt[step+1]
+ } else {
+ token = int64(idx[0]) // choose best token
+ out = append(out, token)
+ }
+
+ cacheValues = outputs[1:]
+ }
+
+ return out[:steps], nil
+}
+
+func forward(model string, token int64, position int64, cacheNames []string, cacheValues []ort.Value) (*ort.Tensor[float32], []string, []ort.Value, error) {
+ inputNames, inputs, _ := initInputs(token, position)
+ outputNames, outputs, logits, _ := initOutputs(position)
+
+ inputNames = append(inputNames, cacheNames...)
+ inputs = append(inputs, cacheValues...)
+
+ session, err := ort.NewAdvancedSession(
+ model,
+ inputNames,
+ outputNames,
+ inputs,
+ outputs,
+ nil,
+ )
+
+ if err != nil {
+ log.Fatal(err)
+ }
+
+ defer session.Destroy()
+
+ if err := session.Run(); err != nil {
+ log.Fatal(err)
+ }
+
+ return logits, outputNames, outputs, nil
+}
+
+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 initInputs(token, position int64) ([]string, []ort.Value, error) {
+ inputNames := []string{"input_ids", "position_ids", "attention_mask"}
+
+ var tokens *ort.Tensor[int64]
+ var positions *ort.Tensor[int64]
+ var attentionMask *ort.Tensor[int64]
+
+ if t, err := ort.NewTensor[int64]([]int64{1, 1}, []int64{token}); err != nil {
+ return nil, nil, err
+ } else {
+ tokens = t
+ }
+
+ if p, err := ort.NewTensor[int64]([]int64{1, 1}, []int64{position}); err != nil {
+ return nil, nil, err
+ } else {
+ positions = p
+ }
+
+ maskData := make([]int64, position+1)
+ maskShape := []int64{1, position + 1}
+
+ for i := range maskData {
+ maskData[i] = 1
+ }
+
+ if m, err := ort.NewTensor[int64](maskShape, maskData); err != nil {
+ return nil, nil, err
+ } else {
+ attentionMask = m
+ }
+
+ inputs := []ort.Value{ort.Value(tokens), ort.Value(positions), ort.Value(attentionMask)}
+
+ return inputNames, inputs, nil
+}
+
+func initOutputs(position int64) ([]string, []ort.Value, *ort.Tensor[float32], error) {
+ outputNames := make([]string, 0, 1+2*nLayers)
+ outputValues := make([]ort.Value, 0, 1+2*nLayers)
+
+ logits, err := ort.NewEmptyTensor[float32]([]int64{1, 1, int64(vocabSize)})
+
+ if err != nil {
+ return nil, nil, nil, err
+ }
+
+ outputNames = append(outputNames, "logits")
+ outputValues = append(outputValues, ort.Value(logits))
+
+ shape := []int64{1, int64(nHeads), position + 1, int64(headDim)}
+
+ for i := range nLayers {
+ kName := fmt.Sprintf("present.%d.key", i)
+ vName := fmt.Sprintf("present.%d.value", i)
+
+ kTensor, _ := ort.NewEmptyTensor[float32](shape)
+ vTensor, _ := ort.NewEmptyTensor[float32](shape)
+
+ outputNames = append(outputNames, kName, vName)
+ outputValues = append(outputValues, ort.Value(kTensor), ort.Value(vTensor))
+ }
+
+ return outputNames, outputValues, logits, nil
+}
+
+func softmax(logits []float32) []float32 {
+ m := logits[0]
+
+ for _, v := range logits {
+ if v > m {
+ m = v
+ }
+ }
+
+ s := float32(0.0)
+ r := make([]float32, len(logits))
+
+ for i, v := range logits {
+ e := float32(math.Exp(float64(v - m)))
+
+ r[i] = e
+ s += e
+ }
+
+ for i := range r {
+ r[i] /= s
+ }
+
+ return r
+}
+
+func topK(p []float32, k int) ([]int, []float32) {
+ n := len(p)
+
+ if k > n {
+ k = n
+ }
+
+ idx := make([]int, n)
+
+ for i := range idx {
+ idx[i] = i
+ }
+
+ sort.Slice(idx, func(i, j int) bool {
+ return p[idx[i]] > p[idx[j]]
+ })
+
+ topIdx := idx[:k]
+ topP := make([]float32, k)
+
+ for i := 0; i < k; i++ {
+ topP[i] = p[topIdx[i]]
+ }
+
+ return topIdx, topP
+}
diff --git a/gpt2/scripts/conv.py b/gpt2/scripts/conv.py
new file mode 100644
index 0000000..cf0cbb8
--- /dev/null
+++ b/gpt2/scripts/conv.py
@@ -0,0 +1,15 @@
+# /// script
+# dependencies = [
+# "torch",
+# "transformers",
+# "optimum",
+# "optimum[onnxruntime]",
+# ]
+# ///
+
+from transformers import AutoTokenizer, AutoModelForCausalLM
+from optimum.onnxruntime import ORTModelForCausalLM
+
+model_id = "gpt2"
+model = ORTModelForCausalLM.from_pretrained(model_id, export=True, use_cache=True)
+model.save_pretrained("onnx-gpt2") \ No newline at end of file