summaryrefslogtreecommitdiff
path: root/research/entropy/experiment.go
diff options
context:
space:
mode:
Diffstat (limited to 'research/entropy/experiment.go')
-rw-r--r--research/entropy/experiment.go145
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)
}