diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-06-03 18:05:03 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-06-03 18:05:03 +0200 |
| commit | 62e116eeb7a25ae78ba490219083e1bdb1ea17c5 (patch) | |
| tree | 1b0ee8344ade8521f60e1db6a33bd2822374bd6c | |
| parent | dd984a2b5beb8bc48e32242e1cfabbfbc5353086 (diff) | |
Add flags to lesci command
| -rw-r--r-- | research/lesci/cmd/lesci/main.go | 61 |
1 files changed, 36 insertions, 25 deletions
diff --git a/research/lesci/cmd/lesci/main.go b/research/lesci/cmd/lesci/main.go index 8a99c18..6f72283 100644 --- a/research/lesci/cmd/lesci/main.go +++ b/research/lesci/cmd/lesci/main.go @@ -1,7 +1,7 @@ package main import ( - "fmt" + "flag" "log" "path" @@ -13,38 +13,35 @@ import ( ) func main() { - if err := gpt2.InitializeEnvironment(); err != nil { - log.Fatal(err) - } + obs := flag.String("tok-obs", "", "observed tokenizer") + ctf := flag.String("tok-ctf", "", "counterfactual tokenizer") - sizes := []int{8192, 16384, 32768, 50256, 100512} + chkpt := flag.String("c", "", "") - for i := range len(sizes) - 1 { - e, m := setup(sizes[i], sizes[len(sizes)-1]) + outPath := flag.String("o", "", "") - if err := m.Init(); err != nil { - log.Fatal(err) - } + dry := flag.Bool("dry", false, "") - if err := e.Run(); err != nil { - log.Fatal(err) - } + flag.Parse() - if err := m.Destroy(); err != nil { - log.Fatal(err) - } + if flag.NArg() < 1 { + log.Fatal("usage: lesci [flags] <model>") } - if err := gpt2.DestroyEnvironment(); err != nil { + if *obs == "" || *ctf == "" { + log.Fatal("tokenizers not specified") + } + + modelPath := flag.Arg(0) + + if err := gpt2.InitializeEnvironment(); err != nil { log.Fatal(err) } -} -func setup(control, treatment int) (*lesci.Experiment, *gpt2.Model) { - a := fmt.Sprintf(shelf.Abs("models/mbpe/gpt2_%d_m000_babylm_v2"), control) - b := fmt.Sprintf(shelf.Abs("models/mbpe/gpt2_%d_m000_babylm_v2"), 100512) + m := must(model(path.Join(shelf.Abs(shelf.Item(modelPath)), *chkpt, "model_eval.onnx"), 50256)) - m := must(model(path.Join(a, "model_eval.onnx"), control)) + a := shelf.Abs(shelf.Item(*obs)) + b := shelf.Abs(shelf.Item(*ctf)) cfg := bpe.Config{ Recover: true, @@ -53,11 +50,25 @@ func setup(control, treatment int) (*lesci.Experiment, *gpt2.Model) { t := must(bpe.NewTokenizerFromFiles(path.Join(a, "vocab.json"), path.Join(a, "merges.txt"), cfg)) c := must(bpe.NewTokenizerFromFiles(path.Join(b, "vocab.json"), path.Join(b, "merges.txt"), cfg)) - d := must(dataset.NewFileReader(shelf.Abs("data/babylm/train_100M"), "*.train")) + d := must(dataset.NewParquetReader(shelf.Abs("data/minipile/test"))) + + o := path.Join(shelf.Abs(shelf.Item(*outPath)), path.Base(shelf.Abs(shelf.Item(modelPath))), *chkpt) + + e := must(lesci.NewExperiment(m, t, c, d, o, 50256, 5000)) + + if err := m.Init(); err != nil { + log.Fatal(err) + } - o := fmt.Sprintf(shelf.Abs("results/lesci/m000/babylm_%d_%d"), control, treatment) + if !*dry { + if err := e.Run(); err != nil { + log.Fatal(err) + } + } - return must(lesci.NewExperiment(m, t, c, d, o, control, 5000)), m + if err := gpt2.DestroyEnvironment(); err != nil { + log.Fatal(err) + } } func must[T any](v T, err error) T { |
