summaryrefslogtreecommitdiff
path: root/research/entropy/experiment.go
blob: 1e4814c7aee309311ab715f172a51b31dbb14e37 (plain)
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
}