diff options
Diffstat (limited to 'gpt2')
| -rw-r--r-- | gpt2/.gitignore | 1 | ||||
| -rw-r--r-- | gpt2/cmd/gpt2/main.go | 18 | ||||
| -rw-r--r-- | gpt2/model.go (renamed from gpt2/main.go) | 46 | ||||
| -rw-r--r-- | gpt2/scripts/conv.py | 2 |
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 |
