summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--research/lesci/cmd/lesci/main.go143
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)])
+}