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)
|
|||
|
|
}
|