diff options
Diffstat (limited to 'research/entropy/cmd/main.go')
| -rw-r--r-- | research/entropy/cmd/main.go | 221 |
1 files changed, 221 insertions, 0 deletions
diff --git a/research/entropy/cmd/main.go b/research/entropy/cmd/main.go new file mode 100644 index 0000000..cba8c33 --- /dev/null +++ b/research/entropy/cmd/main.go @@ -0,0 +1,221 @@ +package main + +import ( + "database/sql" + "fmt" + "log" + + _ "github.com/duckdb/duckdb-go/v2" + "go.jknobloc.com/x/research/entropy" + + "github.com/jonasknobloch/mbpe" + "go.jknobloc.com/x/dataset" + "go.jknobloc.com/x/gpt2" + "go.jknobloc.com/x/shelf" + "go.jknobloc.com/x/tokenizer/bpe" +) + +type NoPreTok struct{} + +func (p *NoPreTok) PreTokenize(s string) []string { + return []string{s} +} + +const CorpusMean = -0.7749100363022372 + +const Threshold = -2.0 + +func main() { + dict := mbpe.NewDict() + + if err := dict.Load(shelf.Abs("results/knobloch/minipile/dict.txt")); err != nil { + log.Fatal(err) + } + + if err := gpt2.InitializeEnvironment(); err != nil { + log.Fatal(err) + } + + // m := model() + // t := tokenizer() + // d := data() + // + // if err := entropy.Context(m, t, d); err != nil { + // log.Fatal(err) + // } + + 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) + + foo := tokenizerMBPE() + + type candidate struct { + src string + ids []int + means []float64 + occurrences int64 + } + + candidates := make([]candidate, 0) + baseline := entropy.NewBaseline() + + for i, c := range dict.Items()[10000:] { + if i > 10 { + break + } + + s := c.Src() + + if len(s) < 10 { + continue + } + + ids := foo.Tokenize(s) + + means, occurrences, err := entropy.MeanLogProbs(ids, db) + + if err != nil { + log.Fatal(err) + } + + if occurrences == 0 { + continue + } + + { + occurrences, err := entropy.LogProbs(ids, db) + + if err != nil { + log.Fatal(err) + } + + fmt.Println(ids) + fmt.Println(occurrences) + } + + candidates = append(candidates, candidate{s, ids, means, occurrences}) + + baseline.Add(means) + } + + for _, c := range candidates { + fmt.Println(c.ids) + fmt.Println(foo.Decoder().Decode([]string{c.src})) + fmt.Println(c.means, c.occurrences) + + fmt.Println("excess ", entropy.Excess(c.means, Threshold)) + fmt.Println("spikes ", entropy.Spikes(c.means)) + fmt.Println("deviations", baseline.Deviations(c.means)) + + fmt.Println() + fmt.Println() + fmt.Println() + } + + if err := gpt2.DestroyEnvironment(); err != nil { + log.Fatal(err) + } +} + +// func must[T any](v T, err error) T { +// if err != nil { +// log.Fatal(err) +// } +// +// return v +// } + +func model() *gpt2.Model { + cfg := gpt2.ConfigDefault() + + cfg.VocabSize = 256 + + opts := gpt2.Options{ + WithCache: false, + WithLogits: false, + WithLogProbs: true, + } + + m := gpt2.NewModel(shelf.Abs("models/mbpe/minipile-byte/gpt2_256_m000_minipile_seed/model_eval.onnx"), cfg, opts) + + if err := m.Init(); err != nil { + log.Fatal(err) + } + + return m +} + +func tokenizer() *bpe.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 + } + + return tok +} + +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 +} + +func data() dataset.Reader { + var miniPile *dataset.ParquetReader + + if r, err := dataset.NewParquetReader(shelf.Abs("data/minipile/validation")); err != nil { + log.Fatal(err) + } else { + miniPile = r + } + + return dataset.NewClampedReader(miniPile, 10) +} + +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 +} |
