diff options
| -rw-r--r-- | research/lesci/cmd/lesci/main.go | 20 |
1 files changed, 16 insertions, 4 deletions
diff --git a/research/lesci/cmd/lesci/main.go b/research/lesci/cmd/lesci/main.go index e9c9a99..80c1d8e 100644 --- a/research/lesci/cmd/lesci/main.go +++ b/research/lesci/cmd/lesci/main.go @@ -7,6 +7,7 @@ import ( "go.jknobloc.com/x/dataset" "go.jknobloc.com/x/gpt2" + "go.jknobloc.com/x/onnx" "go.jknobloc.com/x/research/lesci" "go.jknobloc.com/x/shelf" "go.jknobloc.com/x/tokenizer/bpe" @@ -41,7 +42,7 @@ func main() { log.Fatal(err) } - m := must(model(path.Join(shelf.Abs(shelf.Item(modelPath)), *chkpt, "model_eval.onnx"), 50256)) + m := must(model(path.Join(shelf.Abs(shelf.Item(modelPath)), *chkpt, "model_eval.onnx"))) a := shelf.Abs(shelf.Item(*obs)) b := shelf.Abs(shelf.Item(*ctf)) @@ -68,7 +69,10 @@ func main() { expCfg.ClampRulesBeforeFilter = true - e := must(lesci.NewExperiment(m, t, c, d, o, 50256, 5000, expCfg, expOpts)) + cutoff := len(bpe.Vocab(t)) + window := 5000 + + e := must(lesci.NewExperiment(m, t, c, d, o, cutoff, window, expCfg, expOpts)) if err := m.Init(); err != nil { log.Fatal(err) @@ -93,10 +97,18 @@ func must[T any](v T, err error) T { return v } -func model(name string, vocabSize int) (*gpt2.Model, error) { +func model(name string) (*gpt2.Model, error) { cfg := gpt2.ConfigDefault() - cfg.VocabSize = vocabSize + 1 + // Note that tokenizer vocabulary size and actual model dimensions may differ + // depending on whether special tokens are included in the vocabulary. + // Since we can't rely on vocabulary size, we extract the shape manually. + + if _, shape, err := onnx.ExtractInitializer(name, "transformer.wte.weight"); err != nil { + return nil, err + } else { + cfg.VocabSize = shape[0] + } m := gpt2.NewModel(name, cfg, gpt2.Options{ WithCache: false, |
