Files
tensor/stats/kde.go
T

126 lines
4.3 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package stats
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"sourcedock.dev/petrbalvin/tensor/internal/engine"
)
// Kernel density estimation: the smoothed histogram that turns a
// finite sample into a continuous, everywhere-positive density
// estimate, with the bandwidth doing all the work.
// kdeParallelMinPairs is the point-sample pair count one worker must
// carry before the density sweep splits across goroutines: below it the
// hand-off costs more than the kernel evaluations it would carry.
const kdeParallelMinPairs = 1 << 12
// KernelDensity evaluates the Gaussian-kernel density estimate of the
// sample at every point: each sample contributes a unit-variance
// normal of width bandwidth, and the estimate averages them. A
// non-positive bandwidth asks for Silverman's rule,
// 0.9·min(σ, IQR/1.34)·n^(−1/5), the plug-in width for Gaussian
// truth; a degenerate interquartile range falls back to σ alone. The
// sample and the points must be finite: a NaN or ±Inf entry is
// refused by name rather than spread through the estimate.
func KernelDensity(sample *core.Array, bandwidth float64, points *core.Array) (*core.Array, error) {
const name = "KernelDensity"
if sample.NDim() != 1 || points.NDim() != 1 {
return nil, base.Errf("%s: the sample and the points must be vectors", name)
}
if sample.Dtype() == core.Complex || points.Dtype() == core.Complex {
return nil, base.Errf("%s: complex samples are not supported", name)
}
if err := checkFinite(name, "the sample", sample); err != nil {
return nil, err
}
if err := checkFinite(name, "the points", points); err != nil {
return nil, err
}
n := sample.Len()
if n < 2 {
return nil, base.Errf("%s: the sample needs at least two points, got %d", name, n)
}
if bandwidth <= 0 {
bandwidth = silvermanBandwidth(sample)
}
if math.IsNaN(bandwidth) || bandwidth <= 0 {
return nil, base.Errf("%s: the bandwidth resolved to %g, want a positive width", name, bandwidth)
}
// The sample is read once into a plain slice: the O(n·points) sweep
// below then runs without per-pair accessor dispatch. The values are
// the ones FloatAt returns, so the estimate is unchanged bit for bit.
// A rebased view's payload may run past its own count, so the dense
// path is cut to the visible elements.
sampleVals := rawFloats(sample)
if sampleVals == nil {
sampleVals = make([]float64, n)
for i := range sampleVals {
sampleVals[i] = sample.FloatAt(i)
}
} else {
sampleVals = sampleVals[:n]
}
out := core.New(core.Float, points.Len())
vals := out.RawFloats()
// The points are read once too: a dense float64 payload is walked in
// place, and any other layout is widened into a private slice, so the
// sweep below reads no accessor and every value is the one FloatAt
// returns.
pointVals := rawFloats(points)
if pointVals == nil {
pointVals = make([]float64, points.Len())
for i := range pointVals {
pointVals[i] = points.FloatAt(i)
}
} else {
pointVals = pointVals[:points.Len()]
}
norm := 1 / (float64(n) * bandwidth * math.Sqrt(2*math.Pi))
// An output point is written once by the worker that owns it and its
// sum keeps the sample's own order, so the split of the point range
// never moves a bit: the estimate is identical at any worker count.
density := func(start, end int) {
for i := start; i < end; i++ {
x := pointVals[i]
total := 0.0
for _, s := range sampleVals {
z := (x - s) / bandwidth
total += math.Exp(-0.5 * z * z)
}
vals[i] = total * norm
}
}
engine.ParallelMin(points.Len(), max(1, (kdeParallelMinPairs+n-1)/n), density)
return out, nil
}
// silvermanBandwidth computes the plug-in width from the sample's
// spread: 0.9·min(σ, IQR/1.34)·n^(−1/5), with the interquartile
// fallback to σ when the middle of the sample is degenerate.
func silvermanBandwidth(sample *core.Array) float64 {
n := float64(sample.Len())
sigma, err := Std(sample)
if err != nil {
return 0
}
qArr, err := Quantile(sample, []float64{0.25, 0.75})
if err != nil {
return 0
}
iqr := qArr.FloatAt(1) - qArr.FloatAt(0)
spread := sigma
if iqr > 0 {
spread = math.Min(spread, iqr/1.34)
}
if spread <= 0 {
return 0
}
return 0.9 * spread * math.Pow(n, -0.2)
}