From 97779cc3c9c158b204bf803de57b64a435942f13 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Mon, 13 Jul 2026 18:14:52 +0200 Subject: Infer vocab size and cutoff --- research/lesci/cmd/lesci/main.go | 20 ++++++++++++++++---- 1 file 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, -- cgit v1.2.3