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
|
|||
|
|
}
|