diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-07-14 18:10:37 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-07-14 18:10:37 +0200 |
| commit | bbdaea9ffa0e3421486b6ef25b4c49caa28dcd2b (patch) | |
| tree | 00902cbf5fc9a79d79424123166ee52e2749f246 /research | |
| parent | 97779cc3c9c158b204bf803de57b64a435942f13 (diff) | |
Add flags to control evaluator
Diffstat (limited to 'research')
| -rw-r--r-- | research/lesci/cmd/lesci/main.go | 143 |
1 files changed, 120 insertions, 23 deletions
diff --git a/research/lesci/cmd/lesci/main.go b/research/lesci/cmd/lesci/main.go index 80c1d8e..ed652ee 100644 --- a/research/lesci/cmd/lesci/main.go +++ b/research/lesci/cmd/lesci/main.go @@ -4,65 +4,127 @@ import ( "flag" "log" "path" + "slices" "go.jknobloc.com/x/dataset" "go.jknobloc.com/x/gpt2" + "go.jknobloc.com/x/llm" "go.jknobloc.com/x/onnx" "go.jknobloc.com/x/research/lesci" "go.jknobloc.com/x/shelf" "go.jknobloc.com/x/tokenizer/bpe" ) -func main() { - obs := flag.String("tok-obs", "", "observed tokenizer") - ctf := flag.String("tok-ctf", "", "counterfactual tokenizer") +type Config struct { + TokObs string + TokCtf string + + Chkpt string + + OutPath string + + GoldData string + GoldStep int + + BufferWindow int + BufferStride int + + PadLeft bool + PadRight bool + PadTokenID int64 + + BatchSize int + NumWorkers int + + Dry bool + + Model string +} - chkpt := flag.String("c", "", "") +const PadToken = "<|endoftext|>" - outPath := flag.String("o", "", "") +var cfg Config - goldData := flag.String("gold-data", "", "gold data") - goldStep := flag.Int("gold-step", 0, "gold step") +func init() { + flag.StringVar(&cfg.TokObs, "tok-obs", "", "observed tokenizer") + flag.StringVar(&cfg.TokCtf, "tok-ctf", "", "counterfactual tokenizer") - dry := flag.Bool("dry", false, "") + flag.StringVar(&cfg.Chkpt, "c", "", "checkpoint") + flag.StringVar(&cfg.OutPath, "o", "", "output path") + + flag.StringVar(&cfg.GoldData, "gold-data", "", "gold data") + flag.IntVar(&cfg.GoldStep, "gold-step", 0, "gold step") + + flag.IntVar(&cfg.BufferWindow, "buffer-window", 1024, "buffer window") + flag.IntVar(&cfg.BufferStride, "buffer-stride", 512, "buffer stride") + + flag.BoolVar(&cfg.PadLeft, "pad-left", false, "pad left") + flag.BoolVar(&cfg.PadRight, "pad-right", false, "pad right") + flag.Int64Var(&cfg.PadTokenID, "pad-token-id", -1, "pad token ID") + + flag.IntVar(&cfg.BatchSize, "batch-size", 64, "batch size") + flag.IntVar(&cfg.NumWorkers, "num-workers", 32, "num workers") + + flag.BoolVar(&cfg.Dry, "dry", false, "dry run") +} + +func main() { flag.Parse() if flag.NArg() < 1 { log.Fatal("usage: lesci [flags] <model>") } - if *obs == "" || *ctf == "" { + if cfg.TokObs == "" || cfg.TokCtf == "" { log.Fatal("tokenizers not specified") } - modelPath := flag.Arg(0) + cfg.Model = flag.Arg(0) if err := gpt2.InitializeEnvironment(); err != nil { log.Fatal(err) } - m := must(model(path.Join(shelf.Abs(shelf.Item(modelPath)), *chkpt, "model_eval.onnx"))) + m := must(model(path.Join(shelf.Abs(shelf.Item(cfg.Model)), cfg.Chkpt, "model_eval.onnx"))) - a := shelf.Abs(shelf.Item(*obs)) - b := shelf.Abs(shelf.Item(*ctf)) + a := shelf.Abs(shelf.Item(cfg.TokObs)) + b := shelf.Abs(shelf.Item(cfg.TokCtf)) - cfg := bpe.Config{ + tokCfg := bpe.Config{ Recover: true, } - 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)) + t := must(bpe.NewTokenizerFromFiles(path.Join(a, "vocab.json"), path.Join(a, "merges.txt"), tokCfg)) + c := must(bpe.NewTokenizerFromFiles(path.Join(b, "vocab.json"), path.Join(b, "merges.txt"), tokCfg)) + + if !isPrefix(t, c) { + log.Fatal("counterfactual must be a prefix of observed") + } + + if cfg.PadLeft || cfg.PadRight { + var padTokenID int64 + + if id, ok := tokenID(t, PadToken); ok { + padTokenID = int64(id) + } else { + padTokenID = int64(len(bpe.Vocab(t))) + } + + if padTokenID != cfg.PadTokenID { + log.Fatalf("expected pad token ID %d but got %d", padTokenID, cfg.PadTokenID) + } + } 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) + o := path.Join(shelf.Abs(shelf.Item(cfg.OutPath)), path.Base(shelf.Abs(shelf.Item(cfg.Model))), cfg.Chkpt) expOpts := lesci.Options{ ForceContext: false, ForceExtract: true, - GoldDataName: *goldData, - GoldDataStep: *goldStep, + GoldDataName: cfg.GoldData, + GoldDataStep: cfg.GoldStep, } expCfg := lesci.ConfigDefault() @@ -74,11 +136,24 @@ func main() { e := must(lesci.NewExperiment(m, t, c, d, o, cutoff, window, expCfg, expOpts)) + e.SetTokenBufferConfig(llm.TokenBufferConfig{ + Window: cfg.BufferWindow, + Stride: cfg.BufferStride, + PadLeft: cfg.PadLeft, + PadRight: cfg.PadRight, + PadTokenID: cfg.PadTokenID, + }) + + e.SetEvaluatorConfig(llm.EvaluatorConfig{ + BatchSize: cfg.BatchSize, + NumWorkers: cfg.NumWorkers, + }) + if err := m.Init(); err != nil { log.Fatal(err) } - if !*dry { + if !cfg.Dry { if err := e.Run(); err != nil { log.Fatal(err) } @@ -98,7 +173,7 @@ func must[T any](v T, err error) T { } func model(name string) (*gpt2.Model, error) { - cfg := gpt2.ConfigDefault() + modelCfg := gpt2.ConfigDefault() // Note that tokenizer vocabulary size and actual model dimensions may differ // depending on whether special tokens are included in the vocabulary. @@ -107,10 +182,10 @@ func model(name string) (*gpt2.Model, error) { if _, shape, err := onnx.ExtractInitializer(name, "transformer.wte.weight"); err != nil { return nil, err } else { - cfg.VocabSize = shape[0] + modelCfg.VocabSize = shape[0] } - m := gpt2.NewModel(name, cfg, gpt2.Options{ + m := gpt2.NewModel(name, modelCfg, gpt2.Options{ WithCache: false, WithLogits: false, WithLogProbs: true, @@ -118,3 +193,25 @@ func model(name string) (*gpt2.Model, error) { return m, nil } + +func tokenID(tokenizer *bpe.Tokenizer, eot string) (int, bool) { + vocab := bpe.Vocab(tokenizer) + + for i, t := range vocab { + if t == eot { + return i, true + } + } + + return 0, false +} + +func isPrefix(a, b *bpe.Tokenizer) bool { + vocabA := bpe.Vocab(a) + vocabB := bpe.Vocab(b) + + mergesA := bpe.Merges(a) + mergesB := bpe.Merges(b) + + return slices.Equal(vocabA, vocabB[:len(vocabA)]) && slices.Equal(mergesA, mergesB[:len(mergesA)]) +} |
