diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2025-11-10 15:54:26 +0100 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2025-11-10 15:54:26 +0100 |
| commit | 588596b60cc4f01e82e6efe68aa5c185c8e3f413 (patch) | |
| tree | e825738c56ff2ade2e1cb1f72571274529dc5a32 | |
Initial commit
| -rw-r--r-- | go.mod | 5 | ||||
| -rw-r--r-- | go.sum | 2 | ||||
| -rw-r--r-- | main.go | 109 | ||||
| -rw-r--r-- | scripts/conv.py | 15 |
4 files changed, 131 insertions, 0 deletions
@@ -0,0 +1,5 @@ +module bpc + +go 1.24 + +require github.com/yalue/onnxruntime_go v1.22.0 @@ -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= @@ -0,0 +1,109 @@ +package main + +import ( + "fmt" + "log" + "math" + "sort" + + 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() + + inputNames := []string{"input_ids", "position_ids", "attention_mask"} + outputNames := []string{"logits"} + + 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}) + + 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)}, + nil, + ) + + if err != nil { + log.Fatal(err) + } + + defer session.Destroy() + + if err := session.Run(); err != nil { + log.Fatal(err) + } + + logits := output.GetData() + probs := softmax(logits) + + idx, p := topK(probs, 10) + + for i, t := range idx { + fmt.Printf("%.4f [%d]\n", p[i], t) + } +} + +func softmax(logits []float32) []float32 { + m := float32(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/scripts/conv.py b/scripts/conv.py new file mode 100644 index 0000000..07d63f9 --- /dev/null +++ b/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=False) +model.save_pretrained("onnx-gpt2")
\ No newline at end of file |
