diff options
Diffstat (limited to 'research/entropy/experiment.go')
| -rw-r--r-- | research/entropy/experiment.go | 145 |
1 files changed, 118 insertions, 27 deletions
diff --git a/research/entropy/experiment.go b/research/entropy/experiment.go index 1e4814c..b5efad9 100644 --- a/research/entropy/experiment.go +++ b/research/entropy/experiment.go @@ -3,10 +3,22 @@ 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 { @@ -20,28 +32,36 @@ func Run(db *sql.DB, d *dict.Dict[*Entry]) error { // // TODO handle // } - rows, err := db.Query(`SELECT * FROM context`) + var rows *sql.Rows - if err != nil { + if r, err := db.Query(`SELECT * FROM context`); err != nil { return err + } else { + rows = r } defer rows.Close() - buffer := make([]uint8, 0, BufferSize) + bufferTokens := make([]uint8, 0, BufferSize) bufferLogProbs := make([]float32, 0, BufferSize) doc := -1 - maxTokens := 10000000 - currentToken := 0 + 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 currentToken == maxTokens { - // break + if currentTokens == maxTokens { + break } - currentToken++ + currentTokens++ var uid int var token int @@ -52,22 +72,22 @@ func Run(db *sql.DB, d *dict.Dict[*Entry]) error { return err } - if uid != doc || len(buffer) == BufferSize { - buffer = buffer[:0] + if uid != doc || len(bufferTokens) == BufferSize { + bufferTokens = bufferTokens[:0] bufferLogProbs = bufferLogProbs[:0] doc = uid } - buffer = append(buffer, uint8(token)) + bufferTokens = append(bufferTokens, uint8(token)) bufferLogProbs = append(bufferLogProbs, logProb) - n := len(buffer) + n := len(bufferTokens) // fmt.Print(n) for i := range n { - entry, ok := d.GetBytes(buffer[i:n]) + entry, ok := d.GetBytes(bufferTokens[i:n]) if !ok { continue @@ -79,6 +99,8 @@ func Run(db *sql.DB, d *dict.Dict[*Entry]) error { entry.N++ } + + pb.Add(1) } if err := rows.Err(); err != nil { @@ -90,30 +112,99 @@ func Run(db *sql.DB, d *dict.Dict[*Entry]) error { 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)) - // } + pb.Close() + + for _, e := range d.Values() { + spikes(e) + } + + if err := serializeSegments(d); err != nil { + return err + } return nil } -func spikes(logProbs []float32) []bool { - b := make([]bool, len(logProbs)) +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 i := range len(logProbs) { + 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 } - if d := logProbs[i-1] - logProbs[i]; d > 0 { - b[i] = true // TODO threshold + r[e.Encoded] = segmentsEntry{ + Value: segments, + OK: true, } } - return b + 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) } |
