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 }