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