diff options
| -rw-r--r-- | dict/dict.go | 4 | ||||
| -rw-r--r-- | dict/serialize.go | 55 | ||||
| -rw-r--r-- | research/entropy/cmd/main.go | 221 | ||||
| -rw-r--r-- | research/entropy/cmd/segment/main.go (renamed from research/entropy/cmd/dict3/main.go) | 47 | ||||
| -rw-r--r-- | research/entropy/entry.go | 1 | ||||
| -rw-r--r-- | research/entropy/experiment.go | 145 | ||||
| -rw-r--r-- | research/entropy/segment.go | 346 | ||||
| -rw-r--r-- | research/entropy/segmenter.go | 94 |
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 -} |
