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