summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-06 22:59:50 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-06 23:28:04 +0200
commitb9faf92148c7e6432ff0411783f444b7779fa36a (patch)
tree96904cd4cf053b3d272b5f2b73c1141483e33863
parent58bdf1ac7155f4c5076cf4a077a5f39dfdedbac7 (diff)
Handle environment initialization on package level
-rw-r--r--gpt2/cmd/gpt2/main.go16
-rw-r--r--gpt2/model.go10
-rw-r--r--gpt2/model_test.go14
-rw-r--r--gpt2/onnx.go13
-rw-r--r--llm/cmd/eval/main.go8
5 files changed, 46 insertions, 15 deletions
diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go
index 663bdd8..d854519 100644
--- a/gpt2/cmd/gpt2/main.go
+++ b/gpt2/cmd/gpt2/main.go
@@ -10,12 +10,20 @@ import (
)
func main() {
+ if err := gpt2.InitializeEnvironment(); err != nil {
+ log.Fatal(err)
+ }
+
prompt := []int64{464, 2068, 7586}
generate(prompt) // [-13.483142 -11.277906]
score(prompt) // [-13.48314 -11.277912]
_ = prompt
+
+ if err := gpt2.DestroyEnvironment(); err != nil {
+ log.Fatal(err)
+ }
}
func generate(prompt []int64) {
@@ -35,9 +43,7 @@ func generate(prompt []int64) {
fmt.Println(selectLogProbs(logits[:len(logits)-1], prompt[1:]))
- if err := m.Destroy(); err != nil {
- log.Fatal(err)
- }
+ m.Destroy()
}
func score(prompt []int64) {
@@ -55,9 +61,7 @@ func score(prompt []int64) {
fmt.Println(logProbs)
- if err := m.Destroy(); err != nil {
- log.Fatal(err)
- }
+ m.Destroy()
}
func selectLogProbs(logits [][]float32, tokens []int64) []float32 {
diff --git a/gpt2/model.go b/gpt2/model.go
index f492c56..d0d1c2f 100644
--- a/gpt2/model.go
+++ b/gpt2/model.go
@@ -60,12 +60,6 @@ func IntraOpNumThreads() int {
}
func (m *Model) Init() error {
- ort.SetSharedLibraryPath(SharedLibraryPath())
-
- if err := ort.InitializeEnvironment(); err != nil {
- return err
- }
-
m.allocator = NewAllocator(m.config, m.withCache, m.withLogits, m.withLogProbs)
var options *ort.SessionOptions
@@ -99,10 +93,8 @@ func (m *Model) Init() error {
return nil
}
-func (m *Model) Destroy() error {
+func (m *Model) Destroy() {
m.allocator.Destroy()
-
- return ort.DestroyEnvironment()
}
func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]int64, error) {
diff --git a/gpt2/model_test.go b/gpt2/model_test.go
index d197ccc..0072254 100644
--- a/gpt2/model_test.go
+++ b/gpt2/model_test.go
@@ -9,6 +9,20 @@ import (
"testing"
)
+func TestMain(m *testing.M) {
+ if err := InitializeEnvironment(); err != nil {
+ log.Fatal(err)
+ }
+
+ exit := m.Run()
+
+ if err := DestroyEnvironment(); err != nil {
+ log.Fatal(err)
+ }
+
+ os.Exit(exit)
+}
+
func fromModel() []float32 {
prompt := []int64{464, 2068, 7586, 21831, 18045, 625, 262, 16931, 3290}
diff --git a/gpt2/onnx.go b/gpt2/onnx.go
new file mode 100644
index 0000000..0da2ae2
--- /dev/null
+++ b/gpt2/onnx.go
@@ -0,0 +1,13 @@
+package gpt2
+
+import ort "github.com/yalue/onnxruntime_go"
+
+func InitializeEnvironment() error {
+ ort.SetSharedLibraryPath(SharedLibraryPath())
+
+ return ort.InitializeEnvironment()
+}
+
+func DestroyEnvironment() error {
+ return ort.DestroyEnvironment()
+}
diff --git a/llm/cmd/eval/main.go b/llm/cmd/eval/main.go
index 78a865e..43752d9 100644
--- a/llm/cmd/eval/main.go
+++ b/llm/cmd/eval/main.go
@@ -9,8 +9,16 @@ import (
)
func main() {
+ if err := gpt2.InitializeEnvironment(); err != nil {
+ log.Fatal(err)
+ }
+
perplexity()
// logprobs()
+
+ if err := gpt2.DestroyEnvironment(); err != nil {
+ log.Fatal(err)
+ }
}
func data() *dataset.ParquetReader {