summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-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
-rw-r--r--llm/cmd/eval/main.go2
-rw-r--r--research/lesci/cmd/lesci/main.go6
7 files changed, 30 insertions, 8 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"
diff --git a/llm/cmd/eval/main.go b/llm/cmd/eval/main.go
index 88be500..8b61927 100644
--- a/llm/cmd/eval/main.go
+++ b/llm/cmd/eval/main.go
@@ -40,7 +40,7 @@ func model() *gpt2.Model {
WithLogProbs: true,
}
- m := gpt2.NewModel("gpt2/models/base/model_eval.onnx", "0", gpt2.DefaultConfig(), opts)
+ m := gpt2.NewModel("gpt2/models/base/model_eval.onnx", gpt2.DefaultConfig(), opts)
if err := m.Init(); err != nil {
log.Fatal(err)
diff --git a/research/lesci/cmd/lesci/main.go b/research/lesci/cmd/lesci/main.go
index b485c43..ab0a967 100644
--- a/research/lesci/cmd/lesci/main.go
+++ b/research/lesci/cmd/lesci/main.go
@@ -43,7 +43,7 @@ func setup(control, treatment int) (*lesci.Experiment, *gpt2.Model) {
a := fmt.Sprintf("gpt2/models/onnx_eval/gpt2_%d_m000_babylm_v2", control)
b := fmt.Sprintf("gpt2/models/onnx_eval/gpt2_%d_m000_babylm_v2", 100512)
- m := must(model(path.Join(a, "model_eval.onnx"), "0", control))
+ m := must(model(path.Join(a, "model_eval.onnx"), control))
t := must(bpe.NewTokenizerFromFiles(path.Join(a, "vocab.json"), path.Join(a, "merges.txt")))
c := must(bpe.NewTokenizerFromFiles(path.Join(b, "vocab.json"), path.Join(b, "merges.txt")))
d := must(dataset.NewFileReader("dataset/cmd/dataset/tmp/babylm/train_100M", "*.train"))
@@ -61,12 +61,12 @@ func must[T any](v T, err error) T {
return v
}
-func model(name, device string, vocabSize int) (*gpt2.Model, error) {
+func model(name string, vocabSize int) (*gpt2.Model, error) {
cfg := gpt2.DefaultConfig()
cfg.VocabSize = vocabSize + 1
- m := gpt2.NewModel(name, device, cfg, gpt2.Options{
+ m := gpt2.NewModel(name, cfg, gpt2.Options{
WithCache: false,
WithLogits: false,
WithLogProbs: true,