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.go210
1 files changed, 210 insertions, 0 deletions
diff --git a/research/entropy/experiment.go b/research/entropy/experiment.go
new file mode 100644
index 0000000..b5efad9
--- /dev/null
+++ b/research/entropy/experiment.go
@@ -0,0 +1,210 @@
+package entropy
+
+import (
+ "cmp"
+ "database/sql"
+ "encoding/gob"
+ "os"
+ "time"
+
+ "github.com/jonasknobloch/mbpe"
+
+ "go.jknobloc.com/x/dict"
+ "go.jknobloc.com/x/shelf"
+ "go.jknobloc.com/x/tui"
+)
+
+type segmentsEntry struct {
+ Value []string
+ OK bool
+}
+
+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
+ // }
+
+ var rows *sql.Rows
+
+ if r, err := db.Query(`SELECT * FROM context`); err != nil {
+ return err
+ } else {
+ rows = r
+ }
+
+ defer rows.Close()
+
+ bufferTokens := make([]uint8, 0, BufferSize)
+ bufferLogProbs := make([]float32, 0, BufferSize)
+
+ doc := -1
+
+ maxTokens := 100000
+ currentTokens := 0
+
+ pb := tui.NewProgressBar("Adding LogProbs", 20, 58508191, time.Now())
+
+ pb.Start(1*time.Second, func() int {
+ return pb.Completed()
+ })
+
+ for rows.Next() {
+ if currentTokens == maxTokens {
+ break
+ }
+
+ currentTokens++
+
+ 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(bufferTokens) == BufferSize {
+ bufferTokens = bufferTokens[:0]
+ bufferLogProbs = bufferLogProbs[:0]
+
+ doc = uid
+ }
+
+ bufferTokens = append(bufferTokens, uint8(token))
+ bufferLogProbs = append(bufferLogProbs, logProb)
+
+ n := len(bufferTokens)
+
+ // fmt.Print(n)
+
+ for i := range n {
+ entry, ok := d.GetBytes(bufferTokens[i:n])
+
+ if !ok {
+ continue
+ }
+
+ for j := range entry.LogProbs {
+ entry.LogProbs[j] += bufferLogProbs[j]
+ }
+
+ entry.N++
+ }
+
+ pb.Add(1)
+ }
+
+ 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)
+ })
+
+ pb.Close()
+
+ for _, e := range d.Values() {
+ spikes(e)
+ }
+
+ if err := serializeSegments(d); err != nil {
+ return err
+ }
+
+ return nil
+}
+
+func spikes(e *Entry) {
+ for i := range len(e.LogProbs) {
+ if i == 0 {
+ continue
+ }
+
+ if i == 1 && e.TokenIDs[0] == 220 {
+ continue
+ }
+
+ // TODO threshold
+
+ if d := e.LogProbs[i-1] - e.LogProbs[i]; d > 0 {
+ e.Bounds[i] = true
+ }
+ }
+}
+
+func serializeSegments(d *dict.Dict[*Entry]) error {
+ r := make(map[string]segmentsEntry, d.Len())
+
+ for _, e := range d.Values() {
+ segments := make([]string, 0)
+
+ decoded := decode(e.Encoded)
+
+ i := 0
+
+ for j, b := range e.Bounds {
+ if b {
+ segments = append(segments, encode(decoded[i:j]))
+
+ i = j
+ }
+ }
+
+ segments = append(segments, encode(decoded[i:])) // TODO or just use OK false?
+
+ // debug
+ if i == 0 {
+ continue
+ }
+
+ r[e.Encoded] = segmentsEntry{
+ Value: segments,
+ OK: true,
+ }
+ }
+
+ var file *os.File
+
+ if f, err := os.Create(shelf.Abs("results/entropy/minipile/segments.gob")); err != nil {
+ return err
+ } else {
+ file = f
+ }
+
+ defer file.Close()
+
+ enc := gob.NewEncoder(file)
+
+ return enc.Encode(r)
+}
+
+func encode(s string) string {
+ result := ""
+
+ for _, b := range []byte(s) {
+ result += mbpe.BytesChar[b]
+ }
+
+ return result
+}
+
+func decode(s string) string {
+ result := make([]byte, 0)
+
+ for _, r := range s {
+ result = append(result, mbpe.CharBytes[string(r)])
+ }
+
+ return string(result)
+}