From 321802e7b1f77f684a720d1052befd6d36280875 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Tue, 18 Nov 2025 02:04:41 +0100 Subject: Add perplexity evaluation --- gpt2/cmd/gpt2/main.go | 2 +- gpt2/cuda.go | 36 ++++++++++++++++++++++++++++++++++++ gpt2/model.go | 39 ++++++++++++++++++++++++++------------- 3 files changed, 63 insertions(+), 14 deletions(-) create mode 100644 gpt2/cuda.go (limited to 'gpt2') diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go index 96dabeb..b1e8611 100644 --- a/gpt2/cmd/gpt2/main.go +++ b/gpt2/cmd/gpt2/main.go @@ -9,7 +9,7 @@ import ( func main() { prompt := []int64{464, 2068, 7586, 21831} - m := gpt2.NewModel("models/base/model.onnx") + m := gpt2.NewModel("models/base/model.onnx", "") if err := m.Init(); err != nil { log.Fatal(err) diff --git a/gpt2/cuda.go b/gpt2/cuda.go new file mode 100644 index 0000000..ae3fbcd --- /dev/null +++ b/gpt2/cuda.go @@ -0,0 +1,36 @@ +package gpt2 + +import ort "github.com/yalue/onnxruntime_go" + +func SessionsOptionsWithCUDADeviceID(deviceID string) (*ort.SessionOptions, error) { + var sessionOptions *ort.SessionOptions + var cudaProviderOptions *ort.CUDAProviderOptions + + if s, err := ort.NewSessionOptions(); err != nil { + return nil, err + } else { + sessionOptions = s + } + + if c, err := ort.NewCUDAProviderOptions(); err != nil { + return nil, err + } else { + cudaProviderOptions = c + } + + if err := cudaProviderOptions.Update(map[string]string{ + "device_id": deviceID, + }); err != nil { + return nil, err + } + + if err := sessionOptions.AppendExecutionProviderCUDA(cudaProviderOptions); err != nil { + return nil, err + } + + if err := cudaProviderOptions.Destroy(); err != nil { + return nil, err + } + + return sessionOptions, nil +} diff --git a/gpt2/model.go b/gpt2/model.go index 1b7a9d1..fed1712 100644 --- a/gpt2/model.go +++ b/gpt2/model.go @@ -7,6 +7,7 @@ import ( "log" "math" "os" + "slices" "sort" ort "github.com/yalue/onnxruntime_go" @@ -20,12 +21,14 @@ const ( ) type Model struct { - name string + name string + deviceID string } -func NewModel(name string) *Model { +func NewModel(name, deviceID string) *Model { return &Model{ - name: name, + name: name, + deviceID: deviceID, } } @@ -63,7 +66,7 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in out := make([]int64, 0, steps+1) for step := range context + steps { - _, _, outputs, err := forward(m.name, token, step, cacheNames, cacheValues) + _, _, outputs, err := m.forward(m.name, token, step, cacheNames, cacheValues) if err != nil { return nil, err @@ -96,26 +99,42 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in return out[:steps], nil } -func forward(model string, token int64, position int64, cacheNames []string, cacheValues []ort.Value) (*ort.Tensor[float32], []string, []ort.Value, error) { +func (m *Model) forward(model string, token int64, position int64, cacheNames []string, cacheValues []ort.Value) (*ort.Tensor[float32], []string, []ort.Value, error) { inputNames, inputs, _ := initInputs(token, position) outputNames, outputs, logits, _ := initOutputs(position) inputNames = append(inputNames, cacheNames...) inputs = append(inputs, cacheValues...) + var options *ort.SessionOptions + + if m.deviceID != "" { + if opts, err := SessionsOptionsWithCUDADeviceID(m.deviceID); err != nil { + return nil, nil, nil, err + } else { + options = opts + } + } + session, err := ort.NewAdvancedSession( model, inputNames, outputNames, inputs, outputs, - nil, + options, ) if err != nil { log.Fatal(err) } + if options != nil { + if err := options.Destroy(); err != nil { + return nil, nil, nil, err + } + } + defer session.Destroy() if err := session.Run(); err != nil { @@ -211,13 +230,7 @@ func initOutputs(position int64) ([]string, []ort.Value, *ort.Tensor[float32], e } func softmax(logits []float32) []float32 { - m := logits[0] - - for _, v := range logits { - if v > m { - m = v - } - } + m := slices.Max(logits) s := float32(0.0) r := make([]float32, len(logits)) -- cgit v1.3.1