From 588596b60cc4f01e82e6efe68aa5c185c8e3f413 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Mon, 10 Nov 2025 15:54:26 +0100 Subject: Initial commit --- main.go | 109 ++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 109 insertions(+) create mode 100644 main.go (limited to 'main.go') 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 +} -- cgit v1.3.1