diff options
Diffstat (limited to 'research/entropy/cmd/dict3/main.go')
| -rw-r--r-- | research/entropy/cmd/dict3/main.go | 159 |
1 files changed, 159 insertions, 0 deletions
diff --git a/research/entropy/cmd/dict3/main.go b/research/entropy/cmd/dict3/main.go new file mode 100644 index 0000000..a6b0211 --- /dev/null +++ b/research/entropy/cmd/dict3/main.go @@ -0,0 +1,159 @@ +package main + +import ( + "bufio" + "database/sql" + "fmt" + "io" + "log" + "os" + + _ "github.com/duckdb/duckdb-go/v2" + "github.com/jonasknobloch/mbpe" + "go.jknobloc.com/x/dict" + "go.jknobloc.com/x/research/entropy" + "go.jknobloc.com/x/tokenizer/bpe" + + "go.jknobloc.com/x/shelf" +) + +type NoPreTok struct{} + +func (p *NoPreTok) PreTokenize(s string) []string { + return []string{s} +} + +func main() { + o := shelf.Abs("results/runs-entropy/lesci_minipile_test_256_seed_default/gpt2_256_m000_minipile_seed/lesci.db") + + var db *sql.DB + + if database, err := initDatabase(o); err != nil { + log.Fatal(err) + } else { + db = database + } + + defer db.Close() + + db.SetMaxOpenConns(1) + + var file *os.File + + if f, err := os.Open(shelf.Abs("results/knobloch/minipile/dict.txt")); err != nil { + log.Fatal(err) + } else { + file = f + } + + tok := tokenizerMBPE() + + var dict *dict.Dict[*entropy.Entry] + + if d, err := loadDict(file, tok); err != nil { + log.Fatal(err) + } else { + dict = d + } + + if err := entropy.Run(db, dict); err != nil { + log.Fatal(err) + } +} + +func initDatabase(dsn string) (*sql.DB, error) { + var db *sql.DB + + if database, err := sql.Open("duckdb", dsn); err != nil { + return nil, err + } else { + db = database + } + + if err := db.Ping(); err != nil { + _ = db.Close() + + return nil, err + } + + return db, nil +} + +func loadDict(r io.Reader, tok *mbpe.Tokenizer) (*dict.Dict[*entropy.Entry], error) { + d := dict.NewDict[*entropy.Entry]() + + // TODO we could count non empty lines to pre alloc dict + + scanner := bufio.NewScanner(r) + + for scanner.Scan() { + line := scanner.Text() + + if err := scanner.Err(); err != nil { + return nil, err + } + + var s string + var n int + + if _, err := fmt.Sscanf(line, "%s %d", &s, &n); err != nil { + return nil, err + } + + // TODO see cmd/main.go we currently need mbpe tok ?? + + if n < 100 { + continue + } + + ids := tok.Tokenize(s) + + if len(ids) < 4 { + continue + } + + b := to[int, uint8](ids) // TODO that conversion is fucking ass + + // TODO vs decode each buffer suffix and use s as map key + + d.Set(string(b), &entropy.Entry{ + Encoded: s, // TODO encoded or decoded + TokenIDs: b, + LogProbs: make([]float32, len(ids)), + N: 0, + }) + } + + return d, nil +} + +func to[a bpe.Integer, b bpe.Integer](s []a) []b { + r := make([]b, len(s)) + + // TODO handle overflow + + for i, v := range s { + r[i] = b(v) + } + + return r +} + +func tokenizerMBPE() *mbpe.Tokenizer { + var tok *bpe.Tokenizer + + v := shelf.Abs("tokenizers/minipile/tokenizer_gpt2_256_m000_minipile/vocab.json") + m := shelf.Abs("tokenizers/minipile/tokenizer_gpt2_256_m000_minipile/merges.txt") + + if t, err := bpe.NewTokenizerFromFiles(v, m, bpe.DefaultConfig()); err != nil { + log.Fatal(err) + } else { + tok = t + } + + foo := bpe.MBPE(tok) + + foo.SetPreTokenizer(&NoPreTok{}) + + return foo +} |
