diff options
| -rw-r--r-- | research/lesci/cmd/lesci/main.go | 34 |
1 files changed, 32 insertions, 2 deletions
diff --git a/research/lesci/cmd/lesci/main.go b/research/lesci/cmd/lesci/main.go index ed652ee..c19d17e 100644 --- a/research/lesci/cmd/lesci/main.go +++ b/research/lesci/cmd/lesci/main.go @@ -1,6 +1,7 @@ package main import ( + "errors" "flag" "log" "path" @@ -179,10 +180,10 @@ func model(name string) (*gpt2.Model, error) { // 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 { + if s, err := modelVocabSize(name); err != nil { return nil, err } else { - modelCfg.VocabSize = shape[0] + modelCfg.VocabSize = s } m := gpt2.NewModel(name, modelCfg, gpt2.Options{ @@ -194,6 +195,35 @@ func model(name string) (*gpt2.Model, error) { return m, nil } +func modelVocabSize(name string) (int, error) { + onnxModel := must(onnx.NewModel(name)) + + initializers := []string{ + "transformer.wte.weight", + "model.embed_tokens.weight", + } + + var initializer string + + for _, i := range initializers { + if onnxModel.HasInitializer(i) { + initializer = i + + break + } + } + + if initializer == "" { + return 0, errors.New("no matching initializer") + } + + if _, shape, err := onnxModel.ExtractInitializer(initializer); err != nil { + return 0, err + } else { + return shape[0], nil + } +} + func tokenID(tokenizer *bpe.Tokenizer, eot string) (int, bool) { vocab := bpe.Vocab(tokenizer) |
