126 lines
4.3 KiB
Go
126 lines
4.3 KiB
Go
// 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)
|
||
}
|