summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-06-03 18:05:03 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-06-03 18:05:03 +0200
commit62e116eeb7a25ae78ba490219083e1bdb1ea17c5 (patch)
tree1b0ee8344ade8521f60e1db6a33bd2822374bd6c
parentdd984a2b5beb8bc48e32242e1cfabbfbc5353086 (diff)
Add flags to lesci command
-rw-r--r--research/lesci/cmd/lesci/main.go61
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 {