summaryrefslogtreecommitdiff
path: root/main.go
blob: 6b5bd84e67fbb418f1e02611ec5c660cb050ed45 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
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
}