summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--dict/dict.go4
-rw-r--r--dict/serialize.go55
-rw-r--r--research/entropy/cmd/main.go221
-rw-r--r--research/entropy/cmd/segment/main.go (renamed from research/entropy/cmd/dict3/main.go)47
-rw-r--r--research/entropy/entry.go1
-rw-r--r--research/entropy/experiment.go145
-rw-r--r--research/entropy/segment.go346
-rw-r--r--research/entropy/segmenter.go94
8 files changed, 189 insertions, 724 deletions
diff --git a/dict/dict.go b/dict/dict.go
index bc0f8bd..f153932 100644
--- a/dict/dict.go
+++ b/dict/dict.go
@@ -21,6 +21,10 @@ func NewDict[V any]() *Dict[V] {
}
}
+func (d *Dict[V]) Len() int {
+ return len(d.s)
+}
+
func (d *Dict[V]) Set(key string, val V) {
if i, ok := d.m[key]; ok {
d.s[i].Val = val
diff --git a/dict/serialize.go b/dict/serialize.go
new file mode 100644
index 0000000..da97c06
--- /dev/null
+++ b/dict/serialize.go
@@ -0,0 +1,55 @@
+package dict
+
+import (
+ "bufio"
+ "io"
+ "strings"
+)
+
+func LoadDict[V any](r io.Reader, f func(l string) (string, V)) (*Dict[V], error) {
+ //var d *Dict[V]
+ //
+ //if l, err := linesNotEmpty(r); err != nil {
+ // return nil, err
+ //} else {
+ // d = NewDict[V](l)
+ //}
+
+ d := NewDict[V]()
+
+ scanner := bufio.NewScanner(r)
+
+ for scanner.Scan() {
+ line := scanner.Text()
+
+ if err := scanner.Err(); err != nil {
+ return nil, err
+ }
+
+ k, v := f(line)
+
+ d.Set(k, v)
+ }
+
+ return d, nil
+}
+
+func linesNotEmpty(r io.Reader) (int, error) {
+ scanner := bufio.NewScanner(r)
+
+ count := 0
+
+ for scanner.Scan() {
+ line := scanner.Text()
+
+ if strings.TrimSpace(line) != "" {
+ count++
+ }
+ }
+
+ if err := scanner.Err(); err != nil {
+ return 0, err
+ }
+
+ return count, nil
+}
diff --git a/research/entropy/cmd/main.go b/research/entropy/cmd/main.go
deleted file mode 100644
index cba8c33..0000000
--- a/research/entropy/cmd/main.go
+++ /dev/null
@@ -1,221 +0,0 @@
-package main
-
-import (
- "database/sql"
- "fmt"
- "log"
-
- _ "github.com/duckdb/duckdb-go/v2"
- "go.jknobloc.com/x/research/entropy"
-
- "github.com/jonasknobloch/mbpe"
- "go.jknobloc.com/x/dataset"
- "go.jknobloc.com/x/gpt2"
- "go.jknobloc.com/x/shelf"
- "go.jknobloc.com/x/tokenizer/bpe"
-)
-
-type NoPreTok struct{}
-
-func (p *NoPreTok) PreTokenize(s string) []string {
- return []string{s}
-}
-
-const CorpusMean = -0.7749100363022372
-
-const Threshold = -2.0
-
-func main() {
- dict := mbpe.NewDict()
-
- if err := dict.Load(shelf.Abs("results/knobloch/minipile/dict.txt")); err != nil {
- log.Fatal(err)
- }
-
- if err := gpt2.InitializeEnvironment(); err != nil {
- log.Fatal(err)
- }
-
- // m := model()
- // t := tokenizer()
- // d := data()
- //
- // if err := entropy.Context(m, t, d); err != nil {
- // log.Fatal(err)
- // }
-
- o := shelf.Abs("results/runs-entropy/lesci_minipile_test_256_seed_default/gpt2_256_m000_minipile_seed/lesci.db")
-
- var db *sql.DB
-
- if database, err := initDatabase(o); err != nil {
- log.Fatal(err)
- } else {
- db = database
- }
-
- defer db.Close()
-
- db.SetMaxOpenConns(1)
-
- foo := tokenizerMBPE()
-
- type candidate struct {
- src string
- ids []int
- means []float64
- occurrences int64
- }
-
- candidates := make([]candidate, 0)
- baseline := entropy.NewBaseline()
-
- for i, c := range dict.Items()[10000:] {
- if i > 10 {
- break
- }
-
- s := c.Src()
-
- if len(s) < 10 {
- continue
- }
-
- ids := foo.Tokenize(s)
-
- means, occurrences, err := entropy.MeanLogProbs(ids, db)
-
- if err != nil {
- log.Fatal(err)
- }
-
- if occurrences == 0 {
- continue
- }
-
- {
- occurrences, err := entropy.LogProbs(ids, db)
-
- if err != nil {
- log.Fatal(err)
- }
-
- fmt.Println(ids)
- fmt.Println(occurrences)
- }
-
- candidates = append(candidates, candidate{s, ids, means, occurrences})
-
- baseline.Add(means)
- }
-
- for _, c := range candidates {
- fmt.Println(c.ids)
- fmt.Println(foo.Decoder().Decode([]string{c.src}))
- fmt.Println(c.means, c.occurrences)
-
- fmt.Println("excess ", entropy.Excess(c.means, Threshold))
- fmt.Println("spikes ", entropy.Spikes(c.means))
- fmt.Println("deviations", baseline.Deviations(c.means))
-
- fmt.Println()
- fmt.Println()
- fmt.Println()
- }
-
- if err := gpt2.DestroyEnvironment(); err != nil {
- log.Fatal(err)
- }
-}
-
-// func must[T any](v T, err error) T {
-// if err != nil {
-// log.Fatal(err)
-// }
-//
-// return v
-// }
-
-func model() *gpt2.Model {
- cfg := gpt2.ConfigDefault()
-
- cfg.VocabSize = 256
-
- opts := gpt2.Options{
- WithCache: false,
- WithLogits: false,
- WithLogProbs: true,
- }
-
- m := gpt2.NewModel(shelf.Abs("models/mbpe/minipile-byte/gpt2_256_m000_minipile_seed/model_eval.onnx"), cfg, opts)
-
- if err := m.Init(); err != nil {
- log.Fatal(err)
- }
-
- return m
-}
-
-func tokenizer() *bpe.Tokenizer {
- var tok *bpe.Tokenizer
-
- v := shelf.Abs("tokenizers/minipile/tokenizer_gpt2_256_m000_minipile/vocab.json")
- m := shelf.Abs("tokenizers/minipile/tokenizer_gpt2_256_m000_minipile/merges.txt")
-
- if t, err := bpe.NewTokenizerFromFiles(v, m, bpe.DefaultConfig()); err != nil {
- log.Fatal(err)
- } else {
- tok = t
- }
-
- return tok
-}
-
-func tokenizerMBPE() *mbpe.Tokenizer {
- var tok *bpe.Tokenizer
-
- v := shelf.Abs("tokenizers/minipile/tokenizer_gpt2_256_m000_minipile/vocab.json")
- m := shelf.Abs("tokenizers/minipile/tokenizer_gpt2_256_m000_minipile/merges.txt")
-
- if t, err := bpe.NewTokenizerFromFiles(v, m, bpe.DefaultConfig()); err != nil {
- log.Fatal(err)
- } else {
- tok = t
- }
-
- foo := bpe.MBPE(tok)
-
- foo.SetPreTokenizer(&NoPreTok{})
-
- return foo
-}
-
-func data() dataset.Reader {
- var miniPile *dataset.ParquetReader
-
- if r, err := dataset.NewParquetReader(shelf.Abs("data/minipile/validation")); err != nil {
- log.Fatal(err)
- } else {
- miniPile = r
- }
-
- return dataset.NewClampedReader(miniPile, 10)
-}
-
-func initDatabase(dsn string) (*sql.DB, error) {
- var db *sql.DB
-
- if database, err := sql.Open("duckdb", dsn); err != nil {
- return nil, err
- } else {
- db = database
- }
-
- if err := db.Ping(); err != nil {
- _ = db.Close()
-
- return nil, err
- }
-
- return db, nil
-}
diff --git a/research/entropy/cmd/dict3/main.go b/research/entropy/cmd/segment/main.go
index a6b0211..cc6aea4 100644
--- a/research/entropy/cmd/dict3/main.go
+++ b/research/entropy/cmd/segment/main.go
@@ -1,20 +1,20 @@
package main
import (
- "bufio"
"database/sql"
"fmt"
"io"
"log"
"os"
- _ "github.com/duckdb/duckdb-go/v2"
"github.com/jonasknobloch/mbpe"
+
"go.jknobloc.com/x/dict"
"go.jknobloc.com/x/research/entropy"
+ "go.jknobloc.com/x/shelf"
"go.jknobloc.com/x/tokenizer/bpe"
- "go.jknobloc.com/x/shelf"
+ _ "github.com/duckdb/duckdb-go/v2"
)
type NoPreTok struct{}
@@ -40,7 +40,7 @@ func main() {
var file *os.File
- if f, err := os.Open(shelf.Abs("results/knobloch/minipile/dict.txt")); err != nil {
+ if f, err := os.Open(shelf.Abs("results/entropy/minipile/dict.txt")); err != nil {
log.Fatal(err)
} else {
file = f
@@ -80,51 +80,26 @@ func initDatabase(dsn string) (*sql.DB, error) {
}
func loadDict(r io.Reader, tok *mbpe.Tokenizer) (*dict.Dict[*entropy.Entry], error) {
- d := dict.NewDict[*entropy.Entry]()
-
- // TODO we could count non empty lines to pre alloc dict
-
- scanner := bufio.NewScanner(r)
-
- for scanner.Scan() {
- line := scanner.Text()
-
- if err := scanner.Err(); err != nil {
- return nil, err
- }
-
+ return dict.LoadDict[*entropy.Entry](r, func(l string) (string, *entropy.Entry) {
var s string
var n int
- if _, err := fmt.Sscanf(line, "%s %d", &s, &n); err != nil {
- return nil, err
- }
-
- // TODO see cmd/main.go we currently need mbpe tok ??
-
- if n < 100 {
- continue
+ if _, err := fmt.Sscanf(l, "%s %d", &s, &n); err != nil {
+ panic(err)
}
ids := tok.Tokenize(s)
- if len(ids) < 4 {
- continue
- }
-
b := to[int, uint8](ids) // TODO that conversion is fucking ass
- // TODO vs decode each buffer suffix and use s as map key
-
- d.Set(string(b), &entropy.Entry{
+ return string(b), &entropy.Entry{
Encoded: s, // TODO encoded or decoded
TokenIDs: b,
LogProbs: make([]float32, len(ids)),
+ Bounds: make([]bool, len(ids)+1),
N: 0,
- })
- }
-
- return d, nil
+ }
+ })
}
func to[a bpe.Integer, b bpe.Integer](s []a) []b {
diff --git a/research/entropy/entry.go b/research/entropy/entry.go
index 746e34f..3ad2a26 100644
--- a/research/entropy/entry.go
+++ b/research/entropy/entry.go
@@ -4,5 +4,6 @@ type Entry struct {
Encoded string
TokenIDs []uint8
LogProbs []float32
+ Bounds []bool
N int
}
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)
}
diff --git a/research/entropy/segment.go b/research/entropy/segment.go
deleted file mode 100644
index bdab576..0000000
--- a/research/entropy/segment.go
+++ /dev/null
@@ -1,346 +0,0 @@
-package entropy
-
-import (
- "database/sql"
- "math"
-)
-
-// CREATE TABLE context(uid INTEGER, token INTEGER, logprob FLOAT, pos INTEGER)
-//
-// sqlMeanLogProbs matches an exact token id sequence of any length, bound as a
-// single list parameter, and returns one row per position: the position, the
-// mean logprob over all occurrences, and the occurrence count.
-//
-// seq unrolls the parameter into (position, token) pairs. Joining context on
-// token alone gives every position that matches some element of the sequence;
-// pos - i + 1 is the start a match would have to begin at for that element to
-// line up, so counting rows per (uid, start) says how many elements agree
-// there. A start where all of them agree is an occurrence.
-const sqlMeanLogProbs = `
-WITH seq AS (
- SELECT i, s[i] AS token
- FROM (SELECT ?::INTEGER[] AS s) t, range(1, len(s) + 1) r(i)
-),
-cand AS (
- SELECT s.i AS i,
- c.logprob AS logprob,
- count(*) OVER (PARTITION BY c.uid, c.pos - s.i + 1) AS matched
- FROM context c
- JOIN seq s ON c.token = s.token
-)
-SELECT i, avg(logprob) AS mean, count(*) AS occurrences
-FROM cand
-WHERE matched = (SELECT count(*) FROM seq)
-GROUP BY i
-ORDER BY i`
-
-// sqlLogProbs matches the same sequences as sqlMeanLogProbs but skips the
-// aggregation, returning one row per position per occurrence.
-const sqlLogProbs = `
-WITH seq AS (
- SELECT i, s[i] AS token
- FROM (SELECT ?::INTEGER[] AS s) t, range(1, len(s) + 1) r(i)
-),
-cand AS (
- SELECT c.uid AS uid,
- c.pos - s.i + 1 AS start,
- s.i AS i,
- c.logprob AS logprob,
- count(*) OVER (PARTITION BY c.uid, c.pos - s.i + 1) AS matched
- FROM context c
- JOIN seq s ON c.token = s.token
-)
-SELECT uid, start, i, logprob
-FROM cand
-WHERE matched = (SELECT count(*) FROM seq)
-ORDER BY uid, start, i`
-
-// Occurrence is one match of a sequence, holding the logprob of every position
-// in the context it was found in.
-type Occurrence struct {
- Document int
- Offset int
- Values []float64
-}
-
-// LogProbs returns every occurrence of the exact id sequence separately,
-// leaving the contexts uncombined.
-func LogProbs(ids []int, db *sql.DB) ([]Occurrence, error) {
- seq := make([]int32, len(ids)) // the driver panics on []int
-
- for i, id := range ids {
- seq[i] = int32(id)
- }
-
- rows, err := db.Query(sqlLogProbs, seq)
-
- if err != nil {
- return nil, err
- }
-
- defer rows.Close()
-
- r := make([]Occurrence, 0)
-
- for rows.Next() {
- var document, offset, pos int
- var value float64
-
- if err := rows.Scan(&document, &offset, &pos, &value); err != nil {
- return nil, err
- }
-
- if pos == 1 {
- r = append(r, Occurrence{
- Document: document,
- Offset: offset,
- Values: make([]float64, 0, len(ids)),
- })
- }
-
- if len(r) == 0 {
- continue
- }
-
- last := &r[len(r)-1]
-
- last.Values = append(last.Values, value)
- }
-
- if err := rows.Err(); err != nil {
- return nil, err
- }
-
- return r, nil
-}
-
-// MeanLogProbs returns the mean logprob per position over all occurrences of
-// the exact id sequence, together with the number of occurrences. Occurrences
-// of zero yield a nil slice: the sequence was never seen.
-func MeanLogProbs(ids []int, db *sql.DB) ([]float64, int64, error) {
- seq := make([]int32, len(ids)) // the driver panics on []int
-
- for i, id := range ids {
- seq[i] = int32(id)
- }
-
- rows, err := db.Query(sqlMeanLogProbs, seq)
-
- if err != nil {
- return nil, 0, err
- }
-
- defer rows.Close()
-
- values := make([]float64, 0, len(ids))
-
- var occurrences int64
-
- for rows.Next() {
- var pos int
- var mean float64
- var n int64
-
- if err := rows.Scan(&pos, &mean, &n); err != nil {
- return nil, 0, err
- }
-
- values = append(values, mean)
-
- occurrences = n
- }
-
- if err := rows.Err(); err != nil {
- return nil, 0, err
- }
-
- if len(values) == 0 {
- return nil, 0, nil
- }
-
- return values, occurrences, nil
-}
-
-// Excess weights each position by how far it falls below threshold, in nats.
-// Fires wherever surprisal is high in absolute terms, which in practice means
-// the uncertain zone at the start of a chunk.
-func Excess(values []float64, threshold float64) []float64 {
- r := make([]float64, len(values))
-
- for i, v := range values {
- if d := threshold - v; d > 0 {
- r[i] = d
- }
- }
-
- return r
-}
-
-// Spikes weights each position by how much more surprising it is than the one
-// before it, in nats. Catches boundaries an absolute threshold misses, since
-// surprisal decays inside a unit and a boundary shows up as uncertainty rising
-// again, however low the preceding tail got. Position 0 has no predecessor and
-// is always 0.
-func Spikes(values []float64) []float64 {
- r := make([]float64, len(values))
-
- for i := 1; i < len(values); i++ {
- if d := values[i-1] - values[i]; d > 0 {
- r[i] = d
- }
- }
-
- return r
-}
-
-// Baseline accumulates mean logprobs per position index across sequences.
-// Surprisal decays sharply with position -- onsets are uncertain, tails are
-// forced -- so the raw value at a position says little on its own. Deviations
-// measures how far a sequence departs from what its positions usually look
-// like, which also cancels the downward pull the exact-sequence lookup puts on
-// interior positions.
-type Baseline struct {
- sum []float64
- sumSq []float64
- noise []float64
- counts []int
-}
-
-// MinVarianceShare floors the noise correction: however much sampling noise is
-// estimated, a position keeps at least this share of its observed variance.
-// Without it a position where every candidate agrees would correct to near-zero
-// variance and produce enormous deviations from nothing.
-const MinVarianceShare = 0.1
-
-func NewBaseline() *Baseline {
- return &Baseline{}
-}
-
-func (b *Baseline) grow(n int) {
- for len(b.counts) < n {
- b.sum = append(b.sum, 0)
- b.sumSq = append(b.sumSq, 0)
- b.noise = append(b.noise, 0)
- b.counts = append(b.counts, 0)
- }
-}
-
-func (b *Baseline) Add(values []float64) {
- b.grow(len(values))
-
- for i, v := range values {
- b.sum[i] += v
- b.sumSq[i] += v * v
- b.counts[i]++
- }
-}
-
-// AddCandidate is Add for a candidate whose means were estimated from n
-// occurrences with the given per-position variances across those occurrences.
-// The spread of candidate means overstates how much candidates really differ,
-// by the sampling variance of the means themselves; recording it lets stats
-// subtract it instead of passing the inflation on to every deviation.
-func (b *Baseline) AddCandidate(values, variances []float64, n int) {
- b.Add(values)
-
- if n < 2 {
- return
- }
-
- for i, v := range variances {
- if i >= len(b.noise) {
- break
- }
-
- b.noise[i] += v / float64(n)
- }
-}
-
-func (b *Baseline) stats(pos int) (float64, float64, bool) {
- if pos >= len(b.counts) || b.counts[pos] < 2 {
- return 0, 0, false
- }
-
- n := float64(b.counts[pos])
- mean := b.sum[pos] / n
-
- variance := (b.sumSq[pos] - n*mean*mean) / (n - 1)
-
- if variance <= 0 {
- return mean, 0, false
- }
-
- if corrected := variance - b.noise[pos]/n; corrected > variance*MinVarianceShare {
- variance = corrected
- } else {
- variance = variance * MinVarianceShare
- }
-
- return mean, math.Sqrt(variance), true
-}
-
-// Summarize reduces the occurrences of one candidate to per-position means and
-// variances, the inputs AddCandidate needs.
-func Summarize(occurrences []Occurrence) ([]float64, []float64, int) {
- if len(occurrences) == 0 {
- return nil, nil, 0
- }
-
- n := len(occurrences[0].Values)
-
- means := make([]float64, n)
- variances := make([]float64, n)
-
- for _, o := range occurrences {
- for i, v := range o.Values {
- if i >= n {
- break
- }
-
- means[i] += v
- variances[i] += v * v
- }
- }
-
- for i := range means {
- sum := means[i]
- sumSq := variances[i]
-
- means[i] = sum / float64(len(occurrences))
-
- if len(occurrences) < 2 {
- variances[i] = 0
-
- continue
- }
-
- variances[i] = (sumSq - float64(len(occurrences))*means[i]*means[i]) / float64(len(occurrences)-1)
-
- if variances[i] < 0 {
- variances[i] = 0
- }
- }
-
- return means, variances, len(occurrences)
-}
-
-// Deviations scores each position by how many standard deviations more
-// surprising it is than the same position usually is. Positions at or below
-// their positional mean score 0, as they carry no evidence of a boundary.
-func (b *Baseline) Deviations(values []float64) []float64 {
- r := make([]float64, len(values))
-
- for i, v := range values {
- mean, std, ok := b.stats(i)
-
- if !ok {
- continue
- }
-
- if d := (mean - v) / std; d > 0 {
- r[i] = d
- }
- }
-
- return r
-}
diff --git a/research/entropy/segmenter.go b/research/entropy/segmenter.go
deleted file mode 100644
index 1a15111..0000000
--- a/research/entropy/segmenter.go
+++ /dev/null
@@ -1,94 +0,0 @@
-package entropy
-
-import (
- "database/sql"
-)
-
-type Tokenizer interface {
- Tokenize(string) []int
-}
-
-// Criterion turns per-position mean logprobs into per-position boundary
-// weights, zero wherever there is no evidence of a boundary.
-type Criterion func([]float64) []float64
-
-// ExcessCriterion is Excess bound to a threshold. Weights come out in nats
-// below that threshold, so a Segmenter using it wants a cutoff of 0.
-func ExcessCriterion(threshold float64) Criterion {
- return func(values []float64) []float64 {
- return Excess(values, threshold)
- }
-}
-
-// Segmenter turns surprisal into segments, satisfying the mbpe Segmenter
-// interface so it can stand in for Morfessor without touching the trainer.
-//
-// A boundary is placed wherever the criterion weighs a position above cutoff.
-// Note that cutoff applies to the weights, not to the logprobs: with Spikes it
-// is a minimum rise in nats and belongs above zero, with ExcessCriterion the
-// threshold is already baked in and cutoff belongs at zero.
-type Segmenter struct {
- tokenizer Tokenizer
- db *sql.DB
- criterion Criterion
- cutoff float64
-}
-
-func NewSegmenter(tokenizer Tokenizer, db *sql.DB, criterion Criterion, cutoff float64) *Segmenter {
- return &Segmenter{
- tokenizer: tokenizer,
- db: db,
- criterion: criterion,
- cutoff: cutoff,
- }
-}
-
-// Segment reports the segmentation of compound and whether surprisal was known
-// for it at all. Compounds the model never saw come back unsegmented and not
-// ok, which the trainer already treats as a reason to drop alpha to zero.
-//
-// The leading whitespace marker the trainer strips is put back before the
-// lookup: the logprobs were measured on running text, so the first character
-// of a word is only predictable given the space in front of it.
-func (s *Segmenter) Segment(compound string) ([]string, bool) {
- runes := []rune(compound)
-
- if len(runes) < 2 {
- return []string{compound}, false
- }
-
- ids := s.tokenizer.Tokenize("Ġ" + compound)
-
- if len(ids) != len(runes)+1 {
- return []string{compound}, false // not one token per rune, cannot align
- }
-
- values, occurrences, err := MeanLogProbs(ids, s.db)
-
- if err != nil || occurrences == 0 || len(values) != len(ids) {
- return []string{compound}, false
- }
-
- weights := s.criterion(values)
-
- segments := make([]string, 0, 4)
-
- start := 0
-
- // ids[0] is the whitespace marker and ids[1] the first rune of compound, so
- // a boundary there is the start of the chunk and carries no information.
- // Skipping it also drops the onset spike, which dominates every sequence.
- for i := 2; i < len(ids); i++ {
- if weights[i] <= s.cutoff {
- continue
- }
-
- segments = append(segments, string(runes[start:i-1]))
-
- start = i - 1
- }
-
- segments = append(segments, string(runes[start:]))
-
- return segments, true
-}