diff options
Diffstat (limited to 'research/entropy/segment.go')
| -rw-r--r-- | research/entropy/segment.go | 346 |
1 files changed, 0 insertions, 346 deletions
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 -} |
