1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
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
}
|