diff options
Diffstat (limited to 'research/entropy/segment.go')
| -rw-r--r-- | research/entropy/segment.go | 346 |
1 files changed, 346 insertions, 0 deletions
diff --git a/research/entropy/segment.go b/research/entropy/segment.go new file mode 100644 index 0000000..bdab576 --- /dev/null +++ b/research/entropy/segment.go @@ -0,0 +1,346 @@ +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 +} |
