875 lines
28 KiB
Go
875 lines
28 KiB
Go
// 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
|
||
}
|