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/model.go | 39 +++++++++++++++++++++++++++++++++------ 1 file changed, 33 insertions(+), 6 deletions(-) (limited to 'gpt2/model.go') 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 } } -- cgit v1.3.1