summaryrefslogtreecommitdiff
path: root/research/entropy/cmd/dict3
diff options
context:
space:
mode:
Diffstat (limited to 'research/entropy/cmd/dict3')
-rw-r--r--research/entropy/cmd/dict3/main.go159
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
+}