From eb17c4545f525ac94700f61146cc341323a35a50 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Wed, 11 Feb 2026 16:34:17 +0100 Subject: Allow controlling IntraOpNumThreads --- gpt2/cuda.go | 21 +++++++-------------- 1 file changed, 7 insertions(+), 14 deletions(-) (limited to 'gpt2/cuda.go') diff --git a/gpt2/cuda.go b/gpt2/cuda.go index ae3fbcd..83314e0 100644 --- a/gpt2/cuda.go +++ b/gpt2/cuda.go @@ -2,18 +2,11 @@ package gpt2 import ort "github.com/yalue/onnxruntime_go" -func SessionsOptionsWithCUDADeviceID(deviceID string) (*ort.SessionOptions, error) { - var sessionOptions *ort.SessionOptions +func WithCUDAProvider(options *ort.SessionOptions, deviceID string) error { 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 + return err } else { cudaProviderOptions = c } @@ -21,16 +14,16 @@ func SessionsOptionsWithCUDADeviceID(deviceID string) (*ort.SessionOptions, erro if err := cudaProviderOptions.Update(map[string]string{ "device_id": deviceID, }); err != nil { - return nil, err + return err } - if err := sessionOptions.AppendExecutionProviderCUDA(cudaProviderOptions); err != nil { - return nil, err + if err := options.AppendExecutionProviderCUDA(cudaProviderOptions); err != nil { + return err } if err := cudaProviderOptions.Destroy(); err != nil { - return nil, err + return err } - return sessionOptions, nil + return nil } -- cgit v1.3.1