summaryrefslogtreecommitdiff
path: root/gpt2
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-02-11 16:34:17 +0100
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-02-11 16:34:17 +0100
commiteb17c4545f525ac94700f61146cc341323a35a50 (patch)
treeb334153421f6c7792543296177167ce32c55ed49 /gpt2
parenta2d9cf99a9065c6166222c1e9b333b0e62346829 (diff)
Allow controlling IntraOpNumThreads
Diffstat (limited to 'gpt2')
-rw-r--r--gpt2/cuda.go21
-rw-r--r--gpt2/model.go39
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
}
}