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

875 lines
28 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 signal
import (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
import "sourcedock.dev/petrbalvin/tensor/internal/engine"
import (
"math"
"math/bits"
"sync"
)
// Discrete Fourier transforms. FFT returns the forward transform
// and IFFT the inverse, both as complex arrays from any 1-D input; real
// elements convert. Powers of two run an iterative Cooley-Tukey whose
// stages fuse into radix-4 passes above a small floor (a lone radix-2
// stage remains when the stage count is even); every other length runs
// Bluestein's chirp-z transform over a padded power of two, so no
// length is rejected.
// FFT returns the forward discrete Fourier transform of a 1-D array. An
// empty array is an error, and so is any array above rank one: the
// complex payload shortcut would otherwise flatten a matrix silently
// where the real path refuses. The input is never modified: arrays are
// immutable values, so the transform runs on a private copy.
func FFT(a *core.Array) (*core.Array, error) {
if a.NDim() != 1 {
return nil, base.Errf("FFT: needs a 1-D array, got shape %s", base.ShapeText(a.Shape()))
}
if a.Len() == 0 {
return nil, base.Errf("FFT: an empty array has no FFT")
}
vals, err := a.ComplexValues("FFT")
if err != nil {
return nil, err
}
// ComplexValues shares the payload of a dense complex array and hands
// back a private slice for every other input. The transform works in
// place, so only the shared read needs the copy.
if raw := a.RawComplexes(); len(raw) > 0 && &vals[0] == &raw[0] {
vals = append([]complex128(nil), vals...)
}
transform(vals, -1)
return complexFromArrayMust(vals, []int{len(vals)}), nil
}
// IFFT returns the inverse discrete Fourier transform of a 1-D array,
// scaled by 1/n. An empty array is an error and so is any array above
// rank one. Like FFT it leaves its input untouched.
func IFFT(a *core.Array) (*core.Array, error) {
if a.NDim() != 1 {
return nil, base.Errf("IFFT: needs a 1-D array, got shape %s", base.ShapeText(a.Shape()))
}
if a.Len() == 0 {
return nil, base.Errf("IFFT: an empty array has no FFT")
}
vals, err := a.ComplexValues("IFFT")
if err != nil {
return nil, err
}
// The same alias rule as FFT: only a shared payload is copied, the
// scaling below runs in place.
if raw := a.RawComplexes(); len(raw) > 0 && &vals[0] == &raw[0] {
vals = append([]complex128(nil), vals...)
}
transform(vals, +1)
n := float64(len(vals))
for i := range vals {
vals[i] /= complex(n, 0)
}
return complexFromArrayMust(vals, []int{len(vals)}), nil
}
// transform computes the DFT in place: sign -1 is the forward transform,
// +1 the inverse without scaling.
func transform(vals []complex128, sign float64) {
if isPowerOfTwo(len(vals)) {
fftPow2(vals, sign)
return
}
bluestein(vals, sign)
}
// radix4MinN bounds the transform size below which the plain radix-2
// walk stays: the fused stages earn nothing on a transform that fits
// the first-level cache a few times over, and the small sizes are too
// few butterflies to repay the wider stage code.
const radix4MinN = 1 << 6
// fftPow2 dispatches a power-of-two length between the radix-2 walk
// for small transforms and the fused radix-4 stages above the floor.
func fftPow2(vals []complex128, sign float64) {
if len(vals) < radix4MinN {
fftRadix2(vals, sign)
return
}
fftRadix4(vals, sign)
}
// twiddleCacheMax bounds the transform sizes whose twiddle tables stay
// cached between calls (1 MiB of tables per key at the cap). Bigger
// sizes still build their tables, they just do not keep them.
const twiddleCacheMax = 1 << 16
// twiddle tables are keyed by size and direction sign. A row holds the
// twiddle factors of one butterfly stage: entry k is step multiplied
// into itself k times starting from 1, exactly the value the on-the-fly
// recurrence w *= step produced in butterfly k. Reading the table
// therefore cannot change a single bit of the result; it only spares
// recomputing the recurrence in every parallel block.
type twiddleKey struct {
sign float64
n int
}
var (
twiddleMu sync.RWMutex
twiddleTables = map[twiddleKey][][]complex128{}
)
// twiddlesFor returns the per-stage twiddle rows for a size-n power-of-two
// transform with the given direction sign, building them on first use.
func twiddlesFor(sign float64, n int) [][]complex128 {
key := twiddleKey{sign: sign, n: n}
twiddleMu.RLock()
tables, ok := twiddleTables[key]
twiddleMu.RUnlock()
if ok {
return tables
}
tables = make([][]complex128, bits.Len(uint(n))-1)
for length := 2; length <= n; length <<= 1 {
angle := sign * 2 * math.Pi / float64(length)
sinStep, cosStep := math.Sincos(angle)
step := complex(cosStep, sinStep)
half := length / 2
row := make([]complex128, half)
w := complex(1, 0)
for k := range half {
row[k] = w
w *= step
}
tables[bits.TrailingZeros(uint(length))-1] = row
}
if n <= twiddleCacheMax {
twiddleMu.Lock()
twiddleTables[key] = tables
twiddleMu.Unlock()
}
return tables
}
// bit-reversal swap tables are keyed by size alone: the permutation
// does not depend on the direction sign. A table holds the flat pairs
// (i, j) the in-place reversal loop swaps, i < j, each pair exactly
// once. Applying the pairs reproduces the loop's memory state exactly,
// so reading the table cannot change a single bit of the result.
var (
bitReversalMu sync.RWMutex
bitReversalTables = map[int][]int{}
)
// bitReversalFor returns the flat swap-pair table for the bit-reversal
// permutation of a size-n power-of-two transform, building it on first
// use. Pairs are recorded from the same incremental reversal loop the
// in-place pass ran, so the table is exactly the set of swaps that loop
// performs. Pairs are disjoint, which keeps their application order
// irrelevant. The twiddleCacheMax policy applies: bigger sizes still
// build their tables, they just do not keep them.
func bitReversalFor(n int) []int {
bitReversalMu.RLock()
swaps, ok := bitReversalTables[n]
bitReversalMu.RUnlock()
if ok {
return swaps
}
swaps = make([]int, 0, n)
for i, j := 1, 0; i < n; i++ {
bit := n >> 1
for ; j&bit != 0; bit >>= 1 {
j ^= bit
}
j |= bit
if i < j {
swaps = append(swaps, i, j)
}
}
if n <= twiddleCacheMax {
bitReversalMu.Lock()
bitReversalTables[n] = swaps
bitReversalMu.Unlock()
}
return swaps
}
// fftRadix2 runs iterative Cooley-Tukey with bit reversal.
func fftRadix2(vals []complex128, sign float64) {
n := len(vals)
if n < 2 {
return
}
twiddles := twiddlesFor(sign, n)
// Bit-reversal permutation from the cached swap table.
swaps := bitReversalFor(n)
for p := 0; p < len(swaps); p += 2 {
i, j := swaps[p], swaps[p+1]
vals[i], vals[j] = vals[j], vals[i]
}
// Butterflies over doubling block sizes; every block of a stage
// reads the same cached twiddle row. The blocks of one stage are
// mutually independent, so a wide stage of a large transform splits
// its block range across workers; disjoint blocks touch disjoint
// indices, so the split cannot change the result. Narrow stages and
// small transforms stay on the calling goroutine, where the spawn
// cost would not amortise.
for length := 2; length <= n; length <<= 1 {
row := twiddles[bits.TrailingZeros(uint(length))-1]
blocks := n / length
if n >= parallelMinN && length >= parallelMinBlock && blocks > 1 {
engine.Parallel(blocks, func(bs, be int) {
for b := bs; b < be; b++ {
butterflyStage(vals[b*length:(b+1)*length], row)
}
})
continue
}
butterflyStage(vals, row)
}
}
// parallelMinN bounds the transform size whose stages may split across
// workers, parallelMinBlock the stage width worth splitting. At smaller
// sizes a stage does too little work to repay spawning the workers.
// The floor is low enough that a padded correlation or convolution
// (2·n for an n-point signal) still parallelises its widest stages.
const (
parallelMinN = 1 << 14
parallelMinBlock = 1 << 11
)
// butterflyStage runs one in-place butterfly stage over the whole of
// vals: with half = len(row), every block of 2·half entries combines
// vals[start+k] with vals[start+k+half] through row[k]. It is the
// arithmetic of the original serial stage, unchanged. The two halves
// of a block are cut out as disjoint sub-slices so the tight loop
// indexes without bounds checks.
func butterflyStage(vals, row []complex128) {
half := len(row)
if half == 1 {
// The first stage pairs every even index with its odd
// successor: one flat pass pays no per-block slice setup on
// a single butterfly, which is all a block holds here.
w := row[0]
for i := 1; i < len(vals); i += 2 {
u := vals[i-1]
v := vals[i] * w
vals[i-1] = u + v
vals[i] = u - v
}
return
}
if half == 2 {
// Two butterflies per block: the flat form still wins over
// a slice pair per 4 entries.
w0, w1 := row[0], row[1]
for start := 0; start+3 < len(vals); start += 4 {
u := vals[start]
v := vals[start+2] * w0
vals[start] = u + v
vals[start+2] = u - v
u = vals[start+1]
v = vals[start+3] * w1
vals[start+1] = u + v
vals[start+3] = u - v
}
return
}
for start := 0; start < len(vals); start += 2 * half {
lo := vals[start : start+half : start+half]
hi := vals[start+half : start+2*half : start+2*half]
for k := range half {
u := lo[k]
v := hi[k] * row[k]
lo[k] = u + v
hi[k] = u - v
}
}
}
// radix4Tables caches the fused-stage twiddle rows by size and
// direction sign. One entry covers every fused stage of a size-n
// transform; a stage of block 4·L holds three rows of L entries, the
// powers W^k, W^2k and W^3k of the block twiddle W. Every entry is an
// individual sine/cosine evaluation, so each carries the single
// rounding of its argument instead of the k accumulated roundings a
// recurrence step would leave. The twiddleCacheMax policy applies.
var (
radix4Mu sync.RWMutex
radix4Tables = map[twiddleKey][][3][]complex128{}
)
// radix4TwiddlesFor returns the per-stage twiddle rows of a size-n
// fused transform, one [3][]complex128 per stage block 8, 32, 128, …,
// building them on first use.
func radix4TwiddlesFor(sign float64, n int) [][3][]complex128 {
key := twiddleKey{sign: sign, n: n}
radix4Mu.RLock()
tables, ok := radix4Tables[key]
radix4Mu.RUnlock()
if ok {
return tables
}
tables = nil
for length := 8; length <= n; length *= 4 {
l := length / 4
var rows [3][]complex128
for j := range rows {
row := make([]complex128, l)
for k := range row {
angle := sign * 2 * math.Pi * float64(k*(j+1)) / float64(length)
s, c := math.Sincos(angle)
row[k] = complex(c, s)
}
rows[j] = row
}
tables = append(tables, rows)
}
if n <= twiddleCacheMax {
radix4Mu.Lock()
radix4Tables[key] = tables
radix4Mu.Unlock()
}
return tables
}
// fftRadix4 runs a power-of-two Cooley-Tukey whose stage pairs are
// fused into radix-4 passes: one sweep over the payload does the work
// of two radix-2 sweeps and folds one of the four twiddle multiplies
// a radix-2 pair would spend per butterfly. The input ordering is the
// bit-reversal permutation the radix-2 walk uses, so the cached swap
// tables are shared. The walk opens with the flat length-2 stage and,
// when the stage count is even, closes with a lone radix-2 stage.
func fftRadix4(vals []complex128, sign float64) {
n := len(vals)
swaps := bitReversalFor(n)
for p := 0; p < len(swaps); p += 2 {
i, j := swaps[p], swaps[p+1]
vals[i], vals[j] = vals[j], vals[i]
}
butterflyStage(vals, twiddlesFor(sign, 2)[0])
tables := radix4TwiddlesFor(sign, n)
covered := 2
for ti, length := 0, 8; length <= n; ti, length = ti+1, length*4 {
rows := tables[ti]
covered = length
blocks := n / length
// The blocks of one stage are mutually independent, so a wide
// stage of a large transform splits its block range across
// workers on the same policy the radix-2 walk applies.
if n >= parallelMinN && length >= parallelMinBlock && blocks > 1 {
engine.Parallel(blocks, func(bs, be int) {
for b := bs; b < be; b++ {
radix4Stage(vals[b*length:(b+1)*length], rows, -sign)
}
})
continue
}
radix4Stage(vals, rows, -sign)
}
if covered < n {
butterflyStage(vals, twiddlesFor(sign, n)[bits.TrailingZeros(uint(n))-1])
}
}
// radix4Stage runs one fused radix-4 stage over vals, whose blocks of
// 4·L entries (L = len(tw1)) each hold four L-point transforms. Per
// index k the butterfly forms x0, x1·W^k, x2·W^2k, x3·W^3k over the
// four decimated subsequences and combines them with the exact
// rotation ∓i carrying the odd output slots, an exchange of the pair
// (re, im) with a sign rather than a multiply. q is that sign: +1
// rotates the forward transform's branch by −i, −1 the inverse's by
// +i. The base-2 reversal lays the four transforms down in the order
// ≡0, ≡2, ≡1, ≡3, so the odd decimations read across the middle of
// the block while the outputs land in place.
func radix4Stage(vals []complex128, rows [3][]complex128, q float64) {
tw1, tw2, tw3 := rows[0], rows[1], rows[2]
l := len(tw1)
for start := 0; start < len(vals); start += 4 * l {
s0 := vals[start : start+l : start+l]
s1 := vals[start+l : start+2*l : start+2*l]
s2 := vals[start+2*l : start+3*l : start+3*l]
s3 := vals[start+3*l : start+4*l : start+4*l]
for k := range tw1 {
x0 := s0[k]
x1 := s2[k] * tw1[k]
x2 := s1[k] * tw2[k]
x3 := s3[k] * tw3[k]
e0 := x0 + x2
e1 := x0 - x2
f0 := x1 + x3
f1 := x1 - x3
s0[k] = e0 + f0
s2[k] = e0 - f0
s1[k] = complex(real(e1)+q*imag(f1), imag(e1)-q*real(f1))
s3[k] = complex(real(e1)-q*imag(f1), imag(e1)+q*real(f1))
}
}
}
// bluesteinPlan holds the two pieces of a chirp-z transform that depend
// only on the transform length and the direction sign: the chirp itself
// and the convolution kernel already carried into the frequency domain.
// Both are read-only once published, so any number of transforms may
// share one plan.
type bluesteinPlan struct {
chirp []complex128
kernel []complex128 // the kernel's forward transform
m int
}
// bluestein cache, keyed by size and direction sign like the twiddle
// tables. The plan is a pure function of the key, so serving it from the
// cache cannot change a single bit of the result; it only spares the
// chirp rebuild and one of the three transforms per call. The
// twiddleCacheMax policy applies: bigger sizes still build their plan,
// they just do not keep it.
var (
bluesteinMu sync.RWMutex
bluesteinKeys = map[twiddleKey]*bluesteinPlan{}
)
// bluesteinPlanFor returns the plan for a size-n transform with the
// given direction sign, building it on first use.
func bluesteinPlanFor(sign float64, n int) *bluesteinPlan {
key := twiddleKey{sign: sign, n: n}
bluesteinMu.RLock()
plan, ok := bluesteinKeys[key]
bluesteinMu.RUnlock()
if ok {
return plan
}
chirp := make([]complex128, n)
for k := range n {
k2 := (uint64(k) * uint64(k)) % uint64(2*n)
angle := sign * math.Pi * float64(k2) / float64(n)
sinA, cosA := math.Sincos(angle)
chirp[k] = complex(cosA, sinA)
}
m := 1
for m < 2*n-1 {
m <<= 1
}
// kernel is conj(chirp) for the non-negative offsets and its mirror
// for the wrapped negative ones. The kernel is symmetric:
// c_{−l} = c_l.
kernel := make([]complex128, m)
for i := range n {
kernel[i] = conj(chirp[i])
if i != 0 {
kernel[m-n+i] = conj(chirp[n-i])
}
}
fftPow2(kernel, -1)
plan = &bluesteinPlan{chirp: chirp, kernel: kernel, m: m}
if n <= twiddleCacheMax {
bluesteinMu.Lock()
bluesteinKeys[key] = plan
bluesteinMu.Unlock()
}
return plan
}
// bluestein computes the DFT through the chirp-z transform. With the
// chirp c_l = e^{sign·πi·l²/n}, the identity 2jk = j² + k² − (j−k)² gives
// X_k = c_k · Σ_j (x_j·c_j)·conj(c_{j−k}), a circular convolution that
// runs as a multiplication in the frequency domain of a padded power of
// two m ≥ 2n−1. Squares reduce modulo 2n (the chirp's period) to keep the
// angle small and exact.
func bluestein(vals []complex128, sign float64) {
n := len(vals)
plan := bluesteinPlanFor(sign, n)
chirp, m := plan.chirp, plan.m
// a carries the input over the first n slots, zero padded to m so the
// circular convolution has room for every offset.
a := make([]complex128, m)
for i := range vals {
a[i] = vals[i] * chirp[i]
}
fftPow2(a, -1)
kernel := plan.kernel
for i := range a {
a[i] *= kernel[i]
}
fftPow2(a, +1)
scale := complex(1/float64(m), 0)
for i := range a {
a[i] *= scale
}
for i := range vals {
vals[i] = a[i] * chirp[i]
}
}
// conj returns the complex conjugate.
func conj(c complex128) complex128 {
return complex(real(c), -imag(c))
}
func isPowerOfTwo(n int) bool {
return n > 0 && n&(n-1) == 0
}
// FFT2 returns the 2-D discrete Fourier transform. Input shape
// (H, W), output shape (H, W). Separable: FFT each row, then each
// column. Empty arrays are errors.
func FFT2(a *core.Array) (*core.Array, error) {
return fftND(a, 2, false)
}
// IFFT2 returns the inverse 2-D DFT, scaled by 1/(H·W).
func IFFT2(a *core.Array) (*core.Array, error) {
return fftND(a, 2, true)
}
// FFT3 returns the 3-D DFT over shape (D, H, W).
func FFT3(a *core.Array) (*core.Array, error) {
return fftND(a, 3, false)
}
// IFFT3 returns the inverse 3-D DFT.
func IFFT3(a *core.Array) (*core.Array, error) {
return fftND(a, 3, true)
}
// FFTN returns the N-D DFT along the dimensions in dims (or all
// dimensions if dims is empty); IFFTN over the same dims is its
// inverse. Duplicate entries in dims are applied once per entry: a
// dimension listed twice is transformed twice, the second pass
// running on the result of the first.
func FFTN(a *core.Array, dims []int) (*core.Array, error) {
if len(dims) == 0 {
dims = base.RangeN(a.NDim())
}
cur := a
for _, d := range dims {
// The copy condition mirrors fftND: copy while the array is
// still the caller's input, then chain in place.
out, err := fftAlongDimAny(cur, d, false, cur == a)
if err != nil {
return nil, err
}
cur = out
}
return cur, nil
}
// IFFTN is the inverse N-D DFT, the inverse of FFTN over the same
// dims: it divides by the product of the transformed extents, so
// IFFTN(FFTN(a, dims), dims) restores a.
func IFFTN(a *core.Array, dims []int) (*core.Array, error) {
if len(dims) == 0 {
dims = base.RangeN(a.NDim())
}
cur := a
product := 1
for _, d := range dims {
if d < 0 || d >= cur.NDim() {
return nil, base.Errf("IFFTN: dimension %d out of range for shape %s", d, base.ShapeText(cur.Shape()))
}
// The shape does not change along the chain, so the extents can
// be read before each transform; the scale is their product, not
// the total length, which differs when dims covers a subset.
product *= cur.Shape()[d]
out, err := fftAlongDimAny(cur, d, true, cur == a)
if err != nil {
return nil, err
}
cur = out
}
// Scale by 1/product: each 1-D IFFT is un-scaled (matches the 1-D
// IFFT convention where the scaling is done once at the end).
scale := 1.0 / float64(product)
oc := cur.RawComplexes()
for i := range oc {
oc[i] = complex(real(oc[i])*scale, imag(oc[i])*scale)
}
return cur, nil
}
// RFFT returns the real-input DFT: input is 1-D real-valued, output
// has length n/2+1 complex entries (the non-redundant half of the
// full complex spectrum). For real-valued input the second half is
// the complex conjugate of the first.
func RFFT(a *core.Array) (*core.Array, error) {
if a.NDim() != 1 {
return nil, base.Errf("RFFT: needs a 1-D real input, got shape %s", base.ShapeText(a.Shape()))
}
if a.Dtype() == core.Complex {
return nil, base.Errf("RFFT: input must be real")
}
n := a.Len()
if n == 0 {
return nil, base.Errf("RFFT: an empty array has no FFT")
}
vals := make([]complex128, n)
for i := range n {
vals[i] = complex(a.FloatAt(i), 0)
}
transform(vals, -1)
half := n/2 + 1
out := make([]complex128, half)
copy(out, vals[:half])
return complexFromArrayMust(out, []int{half}), nil
}
// IRFFT is the inverse of RFFT: it takes the non-redundant half of a
// real spectrum and returns the original real signal of length n. n
// must be at least 1; if zero, it defaults to 2·(len−1), or to 1 when
// the spectrum holds the single DC bin.
func IRFFT(a *core.Array, n int) (*core.Array, error) {
if a.NDim() != 1 {
return nil, base.Errf("IRFFT: needs a 1-D complex input, got shape %s", base.ShapeText(a.Shape()))
}
if a.Dtype() != core.Complex {
return nil, base.Errf("IRFFT: input must be complex")
}
half := a.Len()
if half < 1 {
return nil, base.Errf("IRFFT: the spectrum must have at least one entry, got %d", half)
}
if n == 0 {
if half == 1 {
n = 1
} else {
n = 2 * (half - 1)
}
}
if n < 1 {
return nil, base.Errf("IRFFT: n must be at least 1, got %d", n)
}
if half != n/2+1 {
return nil, base.Errf("IRFFT: spectrum length %d doesn't match n=%d", half, n)
}
full := make([]complex128, n)
for i := range half {
full[i] = a.ComplexAt(i)
}
for i := half; i < n; i++ {
full[i] = complex(real(full[n-i]), -imag(full[n-i]))
}
transform(full, +1)
scale := complex(1/float64(n), 0)
out := make([]float64, n)
for i := range n {
out[i] = real(full[i] * scale)
}
return floatsFromArrayMust(out, []int{n}), nil
}
// FFTFreq returns the discrete Fourier transform sample frequencies
// for a signal of length n with sample spacing d. d=1.0 by default.
// Returns the positive and negative frequencies in the standard FFT
// order: [0, 1/n, 2/n, ..., -1/2, ..., -1/n] / d.
func FFTFreq(n int, d float64) *core.Array {
if n <= 0 {
return core.New(core.Float, []int{0}...)
}
if d == 0 {
d = 1
}
vals := make([]float64, n)
// Standard FFT order: 0, 1, …, ⌈n/2⌉−1 over the positive half,
// then −⌊n/2⌋, …, −1; for even n the Nyquist slot n/2 carries
// the negative −1/2, as the doc's "…, −1/2, …" promises.
posHalf := (n-1)/2 + 1
for i := range posHalf {
vals[i] = float64(i) / float64(n) / d
}
for i := posHalf; i < n; i++ {
vals[i] = float64(i-n) / float64(n) / d
}
return floatsFromArrayMust(vals, []int{n})
}
// fftND computes the n-D FFT (rank=2 or 3). Separable over each axis.
func fftND(a *core.Array, rank int, inverse bool) (*core.Array, error) {
if a.NDim() != rank {
return nil, base.Errf("fft%d: needs a %d-D array, got shape %s", rank, rank, base.ShapeText(a.Shape()))
}
if a.Len() == 0 {
return nil, base.Errf("fft%d: an empty array has no transform", rank)
}
cur := a
// Promote to complex on first iteration.
if cur.Dtype() != core.Complex {
vals := make([]complex128, cur.Len())
for i := range vals {
vals[i] = complex(cur.FloatAt(i), 0)
}
cur = complexFromArrayMust(vals, cur.Shape())
}
for d := rank - 1; d >= 0; d-- {
// The first pass copies while cur is still the caller's own
// input (unless the promotion above already built a private
// payload); every later pass reads the previous pass's own
// output, which no one else can see, so it transforms in
// place. Either way the caller's input keeps its bits
// untouched.
out, err := fftAlongDimComplex(cur, d, inverse, cur == a)
if err != nil {
return nil, err
}
cur = out
}
// Scale for inverse: divide by total size.
if inverse {
scale := 1.0 / float64(cur.Len())
oc := cur.RawComplexes()
for i := range oc {
oc[i] = complex(real(oc[i])*scale, imag(oc[i])*scale)
}
}
return cur, nil
}
// fftAlongDimAny is the public-facing entry point: it accepts real
// or complex input and promotes to complex internally. Used by
// FFTN/IFFTN, which need to chain along multiple axes. copyIn may
// transform the payload in place only when the caller owns it; a
// promotion always builds a private payload, so it copies nothing
// regardless.
func fftAlongDimAny(a *core.Array, dim int, inverse, copyIn bool) (*core.Array, error) {
if a.Dtype() != core.Complex {
vals := make([]complex128, a.Len())
for i := range vals {
vals[i] = complex(a.FloatAt(i), 0)
}
a = complexFromArrayMust(vals, a.Shape())
copyIn = false
}
return fftAlongDimComplex(a, dim, inverse, copyIn)
}
// fftLineFloor returns the smallest number of line transforms worth a
// worker's spawn: a line of length n costs about n·⌈log₂n⌉ butterfly
// steps, and a strided line pays a gather and a scatter of the same
// length on top, so the floor is the count carrying a fixed budget of
// them. The stage walk inside a line keeps its own parallelMinBlock
// floor, so a split set of short lines stays serial however many lines
// it holds.
func fftLineFloor(n int, strided bool) int {
steps := n * max(1, bits.Len(uint(n-1)))
if strided {
steps *= 2
}
return workFloorFor(steps)
}
// fftAlongDimComplex applies a 1-D FFT along a single dimension of a
// complex-valued array. Returns a new complex array, or the input
// itself when copyIn is false and the payload is already private: the
// caller must then guarantee no other alias reads the array, because
// the transform runs on it directly. Skipping the copy lets the
// multi-dimension transforms chain without re-copying an array they
// just built.
func fftAlongDimComplex(a *core.Array, dim int, inverse, copyIn bool) (*core.Array, error) {
if dim < 0 || dim >= a.NDim() {
return nil, base.Errf("fft: dimension %d out of range for shape %s", dim, base.ShapeText(a.Shape()))
}
if a.Dtype() != core.Complex {
return nil, base.Errf("fft: input must be complex")
}
nDim := a.Shape()[dim]
if nDim == 0 {
return nil, base.Errf("fft: zero-length dimension %d", dim)
}
stride := 1
for k := dim + 1; k < a.NDim(); k++ {
stride *= a.Shape()[k]
}
blocks := 1
for k := range dim {
blocks *= a.Shape()[k]
}
total := a.Len()
out := a
oc := a.RawComplexes()
if copyIn || oc == nil || a.Strided() {
// The input is the caller's (or not a flat complex payload):
// work on a copy so the caller's array keeps its values. A
// strided payload is gathered in element order; a linear copy
// would keep the physical order instead.
vals := make([]complex128, total)
if a.Strided() {
for i := range total {
vals[i] = a.ComplexAt(i)
}
} else {
copy(vals, oc)
}
out = complexFromArrayMust(vals, a.Shape())
oc = out.RawComplexes()
}
sign := -1.0
if inverse {
sign = +1.0
}
// Every (block, stride) line transforms independently: the lines
// split across workers. Each worker allocates one scratch
// line and reuses it for every position it serves: the line is
// fully rewritten before each transform, so no stale value can
// leak, and two lines never share an output cell. Splitting the
// flat line space (rather than blocks, then strides inside a
// block) keeps single-block layouts, a lone row of columns, as
// parallel as any other.
if stride == 1 {
// Unit stride: every line is already contiguous in the output
// copy, so gather and scatter are pure overhead. Transform each
// line in place: the initial copy above keeps the input
// untouched, and the lines of disjoint blocks never overlap.
engine.ParallelMin(blocks, fftLineFloor(nDim, false), func(bs, be int) {
for b := bs; b < be; b++ {
off := b * nDim
transform(oc[off:off+nDim], sign)
}
})
return out, nil
}
engine.ParallelMin(blocks*stride, fftLineFloor(nDim, true), func(ls, le int) {
slice := make([]complex128, nDim)
for line := ls; line < le; line++ {
b := line / stride
s := line % stride
base := b * nDim * stride
for j := range nDim {
slice[j] = oc[base+j*stride+s]
}
transform(slice, sign)
for j := range nDim {
oc[base+j*stride+s] = slice[j]
}
}
})
return out, nil
}
// complexFromArrayMust and floatsFromArrayMust wrap the taking
// constructors whose lengths always match their shapes by
// construction. They are unexported on purpose: a library that panics
// on a caller's input is a defect, and the callers here cannot fail.
func complexFromArrayMust(vals []complex128, shape []int) *core.Array {
a, err := core.ComplexFromArray(vals, shape...)
if err != nil {
panic("signal: " + err.Error())
}
return a
}
func floatsFromArrayMust(vals []float64, shape []int) *core.Array {
a, err := core.FloatsFromArray(vals, shape...)
if err != nil {
panic("signal: " + err.Error())
}
return a
}