diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-23 20:14:26 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-23 20:19:46 +0200 |
| commit | d35f5fad1d26fc09cfd98f6e7faf17d6a48e4e8f (patch) | |
| tree | 208ca3b822b822ec1c5e236f4a86d8971a79b575 /gpt2 | |
| parent | d3276058dd057178e1218a6b76f2ec42d713d6e8 (diff) | |
Use platform specific defaults for CUDA device ID
Diffstat (limited to 'gpt2')
| -rw-r--r-- | gpt2/cmd/gpt2/main.go | 4 | ||||
| -rw-r--r-- | gpt2/model.go | 14 | ||||
| -rw-r--r-- | gpt2/model_test.go | 2 | ||||
| -rw-r--r-- | gpt2/platform_darwin.go | 5 | ||||
| -rw-r--r-- | gpt2/platform_default.go | 5 |
5 files changed, 26 insertions, 4 deletions
diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go index fc7dec8..06ec535 100644 --- a/gpt2/cmd/gpt2/main.go +++ b/gpt2/cmd/gpt2/main.go @@ -37,7 +37,7 @@ func generate(prompt []int64) { WithLogProbs: false, } - m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_cache.onnx", "0", cfg, opts) + m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_cache.onnx", cfg, opts) if err := m.Init(); err != nil { log.Fatal(err) @@ -69,7 +69,7 @@ func score(prompt []int64) { WithLogProbs: true, } - m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_eval.onnx", "0", cfg, opts) + m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_eval.onnx", cfg, opts) if err := m.Init(); err != nil { log.Fatal(err) diff --git a/gpt2/model.go b/gpt2/model.go index 5fef956..feea423 100644 --- a/gpt2/model.go +++ b/gpt2/model.go @@ -20,7 +20,9 @@ type Model struct { allocator *Allocator } -func NewModel(name string, deviceID string, cfg Config, opts Options) *Model { +func NewModel(name string, cfg Config, opts Options) *Model { + deviceID := CUDADeviceID() + return &Model{ name: name, deviceID: deviceID, @@ -29,6 +31,16 @@ func NewModel(name string, deviceID string, cfg Config, opts Options) *Model { } } +func CUDADeviceID() string { + id, ok := os.LookupEnv("ONNXRUNTIME_CUDA_DEVICE_ID") + + if !ok { + return defaultDeviceID + } + + return id +} + func SharedLibraryPath() string { p, ok := os.LookupEnv("ONNXRUNTIME_SHARED_LIBRARY_PATH") diff --git a/gpt2/model_test.go b/gpt2/model_test.go index e562e0d..051a41a 100644 --- a/gpt2/model_test.go +++ b/gpt2/model_test.go @@ -46,7 +46,7 @@ func model() *Model { WithLogProbs: false, } - m := NewModel("models/base/model.onnx", "0", DefaultConfig(), opts) // TODO check if CUDA is available + m := NewModel("models/base/model.onnx", DefaultConfig(), opts) if err := m.Init(); err != nil { log.Fatal(err) diff --git a/gpt2/platform_darwin.go b/gpt2/platform_darwin.go new file mode 100644 index 0000000..597162e --- /dev/null +++ b/gpt2/platform_darwin.go @@ -0,0 +1,5 @@ +//go:build darwin + +package gpt2 + +const defaultDeviceID = "" diff --git a/gpt2/platform_default.go b/gpt2/platform_default.go new file mode 100644 index 0000000..21b4c14 --- /dev/null +++ b/gpt2/platform_default.go @@ -0,0 +1,5 @@ +//go:build !darwin + +package gpt2 + +const defaultDeviceID = "0" |
