From 588596b60cc4f01e82e6efe68aa5c185c8e3f413 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Mon, 10 Nov 2025 15:54:26 +0100 Subject: Initial commit --- go.mod | 5 +++ go.sum | 2 ++ main.go | 109 ++++++++++++++++++++++++++++++++++++++++++++++++++++++++ scripts/conv.py | 15 ++++++++ 4 files changed, 131 insertions(+) create mode 100644 go.mod create mode 100644 go.sum create mode 100644 main.go create mode 100644 scripts/conv.py diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..77fe489 --- /dev/null +++ b/go.mod @@ -0,0 +1,5 @@ +module bpc + +go 1.24 + +require github.com/yalue/onnxruntime_go v1.22.0 diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..f2c6460 --- /dev/null +++ b/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/main.go b/main.go new file mode 100644 index 0000000..6b5bd84 --- /dev/null +++ b/main.go @@ -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 -- cgit v1.3.1