summaryrefslogtreecommitdiff
path: root/gpt2
diff options
context:
space:
mode:
Diffstat (limited to 'gpt2')
-rw-r--r--gpt2/.gitignore1
-rw-r--r--gpt2/cmd/gpt2/main.go18
-rw-r--r--gpt2/model.go (renamed from gpt2/main.go)46
-rw-r--r--gpt2/scripts/conv.py2
4 files changed, 50 insertions, 17 deletions
diff --git a/gpt2/.gitignore b/gpt2/.gitignore
new file mode 100644
index 0000000..8c6790b
--- /dev/null
+++ b/gpt2/.gitignore
@@ -0,0 +1 @@
+/models \ No newline at end of file
diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go
index ff84f42..96dabeb 100644
--- a/gpt2/cmd/gpt2/main.go
+++ b/gpt2/cmd/gpt2/main.go
@@ -4,24 +4,24 @@ 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")
+ prompt := []int64{464, 2068, 7586, 21831}
+
+ m := gpt2.NewModel("models/base/model.onnx")
- if err := ort.InitializeEnvironment(); err != nil {
+ if err := m.Init(); 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 {
+ if out, err := m.Generate(prompt, 5, nil); err != nil {
log.Fatal(err)
} else {
fmt.Printf("\n%v\n", out)
}
+
+ if err := m.Destroy(); err != nil {
+ log.Fatal(err)
+ }
}
diff --git a/gpt2/main.go b/gpt2/model.go
index ee9e19f..1b7a9d1 100644
--- a/gpt2/main.go
+++ b/gpt2/model.go
@@ -3,8 +3,10 @@ package gpt2
import (
"errors"
"fmt"
+ _ "llm"
"log"
"math"
+ "os"
"sort"
ort "github.com/yalue/onnxruntime_go"
@@ -17,7 +19,37 @@ const (
headDim = 64
)
-func Generate(model string, prompt []int64, steps int64, logits *[][]float32) ([]int64, error) {
+type Model struct {
+ name string
+}
+
+func NewModel(name string) *Model {
+ return &Model{
+ name: name,
+ }
+}
+
+func (m *Model) SharedLibraryPath() string {
+ p, ok := os.LookupEnv("ONNXRUNTIME_SHARED_LIBRARY_PATH")
+
+ if !ok {
+ // TODO embed runtime binaries
+ }
+
+ return p
+}
+
+func (m *Model) Init() error {
+ ort.SetSharedLibraryPath(m.SharedLibraryPath())
+
+ return ort.InitializeEnvironment()
+}
+
+func (m *Model) Destroy() error {
+ return ort.DestroyEnvironment()
+}
+
+func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]int64, error) {
if len(prompt) == 0 {
return nil, errors.New("empty prompt")
}
@@ -31,7 +63,7 @@ func Generate(model string, prompt []int64, steps int64, logits *[][]float32) ([
out := make([]int64, 0, steps+1)
for step := range context + steps {
- _, _, outputs, err := forward(model, token, step, cacheNames, cacheValues)
+ _, _, outputs, err := forward(m.name, token, step, cacheNames, cacheValues)
if err != nil {
return nil, err
@@ -43,13 +75,13 @@ func Generate(model string, prompt []int64, steps int64, logits *[][]float32) ([
*logits = append(*logits, l)
}
- idx, p := topK(softmax(l), 5)
+ idx, _ := topK(softmax(l), 5)
- fmt.Printf("\n%d\n\n", token)
+ // fmt.Printf("\n%d\n\n", token)
- for i, t := range idx {
- fmt.Printf("%.4f %.4f [%d]\n", l[t], p[i], t)
- }
+ // for i, t := range idx {
+ // fmt.Printf("%.4f %.4f [%d]\n", l[t], p[i], t)
+ // }
if step < context-1 {
token = prompt[step+1]
diff --git a/gpt2/scripts/conv.py b/gpt2/scripts/conv.py
index cf0cbb8..e3773e2 100644
--- a/gpt2/scripts/conv.py
+++ b/gpt2/scripts/conv.py
@@ -12,4 +12,4 @@ from optimum.onnxruntime import ORTModelForCausalLM
model_id = "gpt2"
model = ORTModelForCausalLM.from_pretrained(model_id, export=True, use_cache=True)
-model.save_pretrained("onnx-gpt2") \ No newline at end of file
+model.save_pretrained("../models/base") \ No newline at end of file