// Copyright (c) 2026 Petr Balvín (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) }