summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2025-11-13 19:25:27 +0100
committerJonas Knobloch <jonas.knobloch@t-online.de>2025-11-13 20:10:37 +0100
commit019a25a6082a8cde372304e34dcd0ae3d5e875ed (patch)
tree908157de7d20b67d8436a4fbed933c59f4a5d16e
parent148d5910f3577d7f8ed6aef57416a5f9d17efdc6 (diff)
Refactor into workspace
-rw-r--r--go.work3
-rw-r--r--gpt2/cmd/gpt2/main.go27
-rw-r--r--gpt2/go.mod (renamed from go.mod)2
-rw-r--r--gpt2/go.sum (renamed from go.sum)0
-rw-r--r--gpt2/main.go (renamed from main.go)22
-rw-r--r--gpt2/scripts/conv.py (renamed from scripts/conv.py)0
6 files changed, 33 insertions, 21 deletions
diff --git a/go.work b/go.work
new file mode 100644
index 0000000..bddb430
--- /dev/null
+++ b/go.work
@@ -0,0 +1,3 @@
+go 1.24.7
+
+use ./gpt2
diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go
new file mode 100644
index 0000000..ff84f42
--- /dev/null
+++ b/gpt2/cmd/gpt2/main.go
@@ -0,0 +1,27 @@
+package main
+
+import (
+ "fmt"
+ "gpt2"
+ "log"
+
+ 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()
+
+ prompt := []int64{464, 2068, 7586, 21831}
+
+ if out, err := gpt2.Generate("scripts/onnx-gpt2/model.onnx", prompt, 5, nil); err != nil {
+ log.Fatal(err)
+ } else {
+ fmt.Printf("\n%v\n", out)
+ }
+}
diff --git a/go.mod b/gpt2/go.mod
index 77fe489..2ca372e 100644
--- a/go.mod
+++ b/gpt2/go.mod
@@ -1,4 +1,4 @@
-module bpc
+module gpt2
go 1.24
diff --git a/go.sum b/gpt2/go.sum
index f2c6460..f2c6460 100644
--- a/go.sum
+++ b/gpt2/go.sum
diff --git a/main.go b/gpt2/main.go
index 74e7adb..ee9e19f 100644
--- a/main.go
+++ b/gpt2/main.go
@@ -1,4 +1,4 @@
-package main
+package gpt2
import (
"errors"
@@ -17,25 +17,7 @@ const (
headDim = 64
)
-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()
-
- prompt := []int64{464, 2068, 7586, 21831}
-
- if out, err := generate("scripts/onnx-gpt2/model.onnx", prompt, 5, nil); err != nil {
- log.Fatal(err)
- } else {
- fmt.Printf("\n%v\n", out)
- }
-}
-
-func generate(model string, prompt []int64, steps int64, logits *[][]float32) ([]int64, error) {
+func Generate(model string, prompt []int64, steps int64, logits *[][]float32) ([]int64, error) {
if len(prompt) == 0 {
return nil, errors.New("empty prompt")
}
diff --git a/scripts/conv.py b/gpt2/scripts/conv.py
index cf0cbb8..cf0cbb8 100644
--- a/scripts/conv.py
+++ b/gpt2/scripts/conv.py