summaryrefslogtreecommitdiff
path: root/research/entropy/segment.go
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-09-09 02:41:24 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-09-09 02:41:24 +0200
commit3c5f4c4279e6c9d6479e40eff4fbe1add59c1b3a (patch)
treec12ce174c117b0205bc3426b6cc2501f7e709f92 /research/entropy/segment.go
parent8fd3f8aeec85f47d43ef36a81456d48640ca33f9 (diff)
WIP
Diffstat (limited to 'research/entropy/segment.go')
-rw-r--r--research/entropy/segment.go346
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
+}