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 /research | |
| parent | d3276058dd057178e1218a6b76f2ec42d713d6e8 (diff) | |
Use platform specific defaults for CUDA device ID
Diffstat (limited to 'research')
| -rw-r--r-- | research/lesci/cmd/lesci/main.go | 6 |
1 files changed, 3 insertions, 3 deletions
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, |
