summaryrefslogtreecommitdiff
path: root/research
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 /research
parentd3276058dd057178e1218a6b76f2ec42d713d6e8 (diff)
Use platform specific defaults for CUDA device ID
Diffstat (limited to 'research')
-rw-r--r--research/lesci/cmd/lesci/main.go6
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,