Files
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

126 lines
4.3 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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)
}