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