Files
tensor/signal/rankfilter.go
T

259 lines
8.9 KiB
Go
Raw 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 signal
import (
"math"
"slices"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"sourcedock.dev/petrbalvin/tensor/internal/engine"
)
// Rank-statistic filters. Every output sample is one order
// statistic of the window around it: the median for the middle rank,
// the minimum for rank 0, the maximum for the last rank. Order
// statistics ignore the shape of the tail inside the window, which is
// why the median removes an impulsive spike outright where a linear
// filter would smear it, and why a rank filter holds an edge where a
// moving average would round it off.
//
// The kernels hold two conventions fixed:
//
// - Edges. Windows near a boundary are truncated to the samples
// that exist (the SavitzkyGolay convention), so the output is
// complete without padding, and the requested rank scales to the
// truncation: rank k of a full window becomes rank k·m/window of
// an m-sample truncation, so a median stays the middle of
// whatever samples the edge left and an extremum stays an
// extremum.
// - NaN. As in the pooling kernels, a NaN in a window propagates:
// the window's order statistic answers NaN, rather than sorting
// it silently into some position.
//
// Non-complex inputs widen once through widenFloats, so the values
// compared are exactly the elements the accessors report; the outputs
// are float64, as for the pooling kernels.
// rankMinPointsPerWorker is the smallest per-worker chunk of output
// points the 1-D rank sweep splits for: below it a chunk's window
// sorts no longer pay the worker spawn cost.
const rankMinPointsPerWorker = 1 << 10
// MedianFilter returns the running median of the rank-1 signal x
// over an odd window: each sample becomes the middle order statistic
// of its neighbourhood. A polynomial-free smoother that removes
// isolated spikes whole and holds monotone ramps and edges still.
// The window must be odd and at least 3, must not exceed the signal,
// and a NaN in a window propagates into that output sample.
func MedianFilter(x *core.Array, window int) (*core.Array, error) {
return RankFilter(x, window, window/2)
}
// RankFilter returns the k-th order statistic (ascending, 0-based) of
// each window of the rank-1 signal x: rank 0 is the running minimum,
// window−1 the running maximum, window/2 the median MedianFilter
// takes. The window must be odd and at least 3 and must not exceed
// the signal; k must lie in [0, window).
func RankFilter(x *core.Array, window, k int) (*core.Array, error) {
const name = "RankFilter"
if err := rankGates(name, x, 1, window, k, window); err != nil {
return nil, err
}
n := x.Len()
if window > n {
return nil, base.Errf("%s: window %d exceeds the signal length %d", name, window, n)
}
src := widenFloats(x)
out := core.New(core.Float, n)
dst := out.RawFloats()
// Output points split across workers: each owns disjoint slots
// and its own window scratch, so the split cannot move a bit.
engine.ParallelMin(n, rankMinPointsPerWorker, func(ps, pe int) {
win := make([]float64, window)
// sorted mirrors the current window in ascending order, so the
// rank's element is a direct read instead of a fresh sort of the
// same values. Neighbouring windows share all but one element, so
// the mirror slides: one removal and one insertion per output
// point.
//
// The mirror answers the rank only while the window holds no zero.
// A zero may sit next to a −0, its equal under comparison but not
// in bits, and which of the two the sort leaves at the rank is the
// sort's own tie choice, which no multiset can predict. Those
// windows take the exact path below, the sort this filter has
// always run. A NaN entering or leaving the window likewise drops
// the mirror, which is rebuilt from the window on the next point.
sorted := make([]float64, 0, window)
zeros := 0
mirror := false
half := window / 2
prevLo, prevHi := 0, 0
for i := ps; i < pe; i++ {
lo := max(i-half, 0)
hi := min(i+half+1, n)
if i > ps && mirror {
// Leave behind what the window no longer covers and take
// up what it now does. Both bounds advance by at most one
// element, so at most one leaves and one enters.
drop := false
for j := prevLo; j < lo && !drop; j++ {
if v := src[j]; !math.IsNaN(v) {
if v == 0 {
zeros--
}
pos, _ := slices.BinarySearch(sorted, v)
sorted = slices.Delete(sorted, pos, pos+1)
continue
}
drop = true
}
for j := prevHi; j < hi && !drop; j++ {
if v := src[j]; !math.IsNaN(v) {
if v == 0 {
zeros++
}
pos, _ := slices.BinarySearch(sorted, v)
sorted = slices.Insert(sorted, pos, v)
continue
}
drop = true
}
mirror = !drop
}
prevLo, prevHi = lo, hi
if !mirror {
// A NaN is not part of the mirror, so a window holding one
// answers NaN and leaves the mirror to be rebuilt later.
sorted, zeros = sorted[:0], 0
nan := false
for j := lo; j < hi; j++ {
v := src[j]
if math.IsNaN(v) {
nan = true
break
}
if v == 0 {
zeros++
}
sorted = append(sorted, v)
}
if nan {
dst[i] = math.NaN()
continue
}
slices.Sort(sorted)
mirror = true
}
count := hi - lo
// The rank scales to the truncation; k below the full
// window's capacity keeps the product below count and
// the index inside the collected prefix.
rank := k * count / window
if zeros > 0 {
for j := lo; j < hi; j++ {
win[j-lo] = src[j]
}
slices.Sort(win[:count])
dst[i] = win[rank]
continue
}
dst[i] = sorted[rank]
}
})
return out, nil
}
// MedianFilter2D returns the running median of the rank-2 image img
// over a square odd window, the standard impulse-noise cleaner: each
// pixel becomes the middle of the values in its neighbourhood, so
// salt-and-pepper dots vanish while steps between regions keep their
// corners. The window must be odd and at least 3 and must fit both
// dimensions; a NaN in a window propagates into that output pixel.
func MedianFilter2D(img *core.Array, window int) (*core.Array, error) {
return RankFilter2D(img, window, window*window/2)
}
// RankFilter2D returns the k-th order statistic (ascending, 0-based)
// of each square window of the rank-2 image img: rank 0 is the
// running minimum (erosion), window²−1 the running maximum (dilation),
// the middle rank the median. The window must be odd and at least 3
// and must fit both dimensions; k must lie in [0, window²).
func RankFilter2D(img *core.Array, window, k int) (*core.Array, error) {
const name = "RankFilter2D"
if err := rankGates(name, img, 2, window, k, window*window); err != nil {
return nil, err
}
h, w := img.Shape()[0], img.Shape()[1]
if window > h || window > w {
return nil, base.Errf("%s: window %d does not fit the %dx%d image", name, window, h, w)
}
src := widenFloats(img)
out := core.New(core.Float, h, w)
dst := out.RawFloats()
// One work item per image row; rows own disjoint output pixels
// and their own window scratch, and every window walks its
// elements in row-major order.
engine.ParallelMin(h, workFloorFor(w*window*window), func(ys, ye int) {
win := make([]float64, window*window)
half := window / 2
for y := ys; y < ye; y++ {
yLo := max(y-half, 0)
yHi := min(y+half+1, h)
for x := range w {
xLo := max(x-half, 0)
xHi := min(x+half+1, w)
count := 0
nan := false
for yy := yLo; yy < yHi && !nan; yy++ {
row := yy * w
for xx := xLo; xx < xHi; xx++ {
v := src[row+xx]
if math.IsNaN(v) {
nan = true
break
}
win[count] = v
count++
}
}
if nan {
dst[y*w+x] = math.NaN()
continue
}
slices.Sort(win[:count])
// The rank scales to the truncation; k below the
// full window's capacity window² keeps the
// product below count and the index inside the
// collected prefix.
dst[y*w+x] = win[k*count/(window*window)]
}
}
})
return out, nil
}
// rankGates checks the contract the four public rank filters share:
// the promised rank, a real dtype, a non-empty input, an odd window
// of at least 3, and a rank inside the window's capacity kMax.
func rankGates(name string, a *core.Array, ndim, window, k, kMax int) error {
if a.NDim() != ndim {
return base.Errf("%s: needs a rank-%d input, got shape %s", name, ndim, base.ShapeText(a.Shape()))
}
if a.Dtype() == core.Complex {
return base.Errf("%s: complex arrays are not supported", name)
}
if a.Len() == 0 {
return base.Errf("%s: the input must not be empty", name)
}
if window < 3 || window%2 == 0 {
return base.Errf("%s: window must be an odd number ≥ 3, got %d", name, window)
}
if k < 0 || k >= kMax {
return base.Errf("%s: k must lie in [0, %d) for window %d, got %d", name, kMax, window, k)
}
return nil
}