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/model.go | |
| parent | a2d9cf99a9065c6166222c1e9b333b0e62346829 (diff) | |
Allow controlling IntraOpNumThreads
Diffstat (limited to 'gpt2/model.go')
| -rw-r--r-- | gpt2/model.go | 39 |
1 files changed, 33 insertions, 6 deletions
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 } } |
