summaryrefslogtreecommitdiff
path: root/research
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-07-13 18:14:52 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-07-13 18:14:52 +0200
commit97779cc3c9c158b204bf803de57b64a435942f13 (patch)
tree65bb0ca4c400954c16874ffa5432523eb4712729 /research
parentbb47fe32789022afc514fca650d08580aa8d22f0 (diff)
Infer vocab size and cutoff
Diffstat (limited to 'research')
-rw-r--r--research/lesci/cmd/lesci/main.go20
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,