diff options
| -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 | ||||
| -rw-r--r-- | llm/cmd/eval/main.go | 2 | ||||
| -rw-r--r-- | research/lesci/cmd/lesci/main.go | 6 |
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, |
