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