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