summaryrefslogtreecommitdiff
path: root/gpt2
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-23 20:14:26 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-23 20:19:46 +0200
commitd35f5fad1d26fc09cfd98f6e7faf17d6a48e4e8f (patch)
tree208ca3b822b822ec1c5e236f4a86d8971a79b575 /gpt2
parentd3276058dd057178e1218a6b76f2ec42d713d6e8 (diff)
Use platform specific defaults for CUDA device ID
Diffstat (limited to 'gpt2')
-rw-r--r--gpt2/cmd/gpt2/main.go4
-rw-r--r--gpt2/model.go14
-rw-r--r--gpt2/model_test.go2
-rw-r--r--gpt2/platform_darwin.go5
-rw-r--r--gpt2/platform_default.go5
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"