diff options
Diffstat (limited to 'research/entropy/experiment.go')
| -rw-r--r-- | research/entropy/experiment.go | 210 |
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) +} |
