summaryrefslogtreecommitdiff
path: root/research/entropy/experiment.go
diff options
context:
space:
mode:
Diffstat (limited to 'research/entropy/experiment.go')
-rw-r--r--research/entropy/experiment.go119
1 files changed, 119 insertions, 0 deletions
diff --git a/research/entropy/experiment.go b/research/entropy/experiment.go
new file mode 100644
index 0000000..1e4814c
--- /dev/null
+++ b/research/entropy/experiment.go
@@ -0,0 +1,119 @@
+package entropy
+
+import (
+ "cmp"
+ "database/sql"
+
+ "go.jknobloc.com/x/dict"
+)
+
+const BufferSize = 1024
+
+func Run(db *sql.DB, d *dict.Dict[*Entry]) error {
+ // m := 0
+ //
+ // for _, v := range d.Values() {
+ // m = max(m, len(v.TokenIDs))
+ // }
+ //
+ // if m > BufferSize {
+ // // TODO handle
+ // }
+
+ rows, err := db.Query(`SELECT * FROM context`)
+
+ if err != nil {
+ return err
+ }
+
+ defer rows.Close()
+
+ buffer := make([]uint8, 0, BufferSize)
+ bufferLogProbs := make([]float32, 0, BufferSize)
+
+ doc := -1
+
+ maxTokens := 10000000
+ currentToken := 0
+
+ for rows.Next() {
+ if currentToken == maxTokens {
+ // break
+ }
+
+ currentToken++
+
+ var uid int
+ var token int
+ var logProb float32
+ var pos int
+
+ if err := rows.Scan(&uid, &token, &logProb, &pos); err != nil {
+ return err
+ }
+
+ if uid != doc || len(buffer) == BufferSize {
+ buffer = buffer[:0]
+ bufferLogProbs = bufferLogProbs[:0]
+
+ doc = uid
+ }
+
+ buffer = append(buffer, uint8(token))
+ bufferLogProbs = append(bufferLogProbs, logProb)
+
+ n := len(buffer)
+
+ // fmt.Print(n)
+
+ for i := range n {
+ entry, ok := d.GetBytes(buffer[i:n])
+
+ if !ok {
+ continue
+ }
+
+ for j := range entry.LogProbs {
+ entry.LogProbs[j] += bufferLogProbs[j]
+ }
+
+ entry.N++
+ }
+ }
+
+ if err := rows.Err(); err != nil {
+ return err
+ }
+
+ // debug
+ d.Sort(func(a, b dict.Entry[*Entry]) int {
+ return cmp.Compare(b.Val.N, a.Val.N)
+ })
+
+ // for _, e := range d.Values() {
+ // if e.N < 100 {
+ // break
+ // }
+ //
+ // fmt.Println(e.Encoded)
+ // fmt.Println(spikes(e.LogProbs))
+ // }
+
+ return nil
+}
+
+func spikes(logProbs []float32) []bool {
+ b := make([]bool, len(logProbs))
+
+ for i := range len(logProbs) {
+ if i == 0 {
+ continue
+ }
+
+ if d := logProbs[i-1] - logProbs[i]; d > 0 {
+ b[i] = true // TODO threshold
+ }
+ }
+
+ return b
+}