diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-02-11 16:34:17 +0100 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-02-11 16:34:17 +0100 |
| commit | eb17c4545f525ac94700f61146cc341323a35a50 (patch) | |
| tree | b334153421f6c7792543296177167ce32c55ed49 /gpt2 | |
| parent | a2d9cf99a9065c6166222c1e9b333b0e62346829 (diff) | |
Allow controlling IntraOpNumThreads
Diffstat (limited to 'gpt2')
| -rw-r--r-- | gpt2/cuda.go | 21 | ||||
| -rw-r--r-- | gpt2/model.go | 39 |
2 files changed, 40 insertions, 20 deletions
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 } diff --git a/gpt2/model.go b/gpt2/model.go index 2dcb452..3655b8a 100644 --- a/gpt2/model.go +++ b/gpt2/model.go @@ -8,6 +8,7 @@ import ( "os" "slices" "sort" + "strconv" ort "github.com/yalue/onnxruntime_go" ) @@ -34,7 +35,7 @@ func NewModel(name, deviceID string) *Model { } } -func (m *Model) SharedLibraryPath() string { +func SharedLibraryPath() string { p, ok := os.LookupEnv("ONNXRUNTIME_SHARED_LIBRARY_PATH") if !ok { @@ -44,8 +45,24 @@ func (m *Model) SharedLibraryPath() string { return p } +func IntraOpNumThreads() int { + s, ok := os.LookupEnv("ONNXRUNTIME_INTRA_OP_NUM_THREADS") + + if !ok { + return 0 + } + + n, err := strconv.Atoi(s) + + if err != nil { + return 0 + } + + return n +} + func (m *Model) Init() error { - ort.SetSharedLibraryPath(m.SharedLibraryPath()) + ort.SetSharedLibraryPath(SharedLibraryPath()) if err := ort.InitializeEnvironment(); err != nil { return err @@ -67,13 +84,23 @@ func (m *Model) Init() error { var options *ort.SessionOptions + if o, err := ort.NewSessionOptions(); err != nil { + return err + } else { + options = o + + defer options.Destroy() + } + if m.deviceID != "" { - if opts, err := SessionsOptionsWithCUDADeviceID(m.deviceID); err != nil { + if err := WithCUDAProvider(options, m.deviceID); err != nil { return err - } else { - options = opts + } + } - defer options.Destroy() + if n := IntraOpNumThreads(); n > 0 { + if err := options.SetIntraOpNumThreads(n); err != nil { + return err } } |
