// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package signal import "sourcedock.dev/petrbalvin/tensor/internal/core" import ( "math" "strings" "testing" ) func TestFFT2(t *testing.T) { // 2×2 identity-like: shift. a := mustFromFloats(t, []float64{ 1, 0, 0, 0, }, 2, 2) got, err := FFT2(a) if err != nil { t.Fatal(err) } if got.Shape()[0] != 2 || got.Shape()[1] != 2 { t.Errorf("FFT2 shape: %v", got.Shape()) } if got.Dtype() != core.Complex { t.Errorf("FFT2 dtype: %s", got.Dtype()) } // Round-trip: IFFT2(FFT2(x)) ≈ x. back, err := IFFT2(got) if err != nil { t.Fatal(err) } for i := range 4 { v, _ := core.ComplexAt(back, i/2, i%2) orig, _ := core.FloatAt(a, i/2, i%2) if math.Abs(real(v)-orig) > 1e-9 { t.Errorf("FFT2 round-trip [%d]: got %v, want %v", i, real(v), orig) } } } func TestFFT3(t *testing.T) { a := mustFromFloats(t, []float64{ 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, }, 3, 2, 2) got, err := FFT3(a) if err != nil { t.Fatal(err) } if got.Shape()[0] != 3 || got.Shape()[1] != 2 || got.Shape()[2] != 2 { t.Errorf("FFT3 shape: %v", got.Shape()) } // Round-trip. back, err := IFFT3(got) if err != nil { t.Fatal(err) } v0, _ := core.ComplexAt(back, 0, 0, 0) if math.Abs(real(v0)-1) > 1e-9 { t.Errorf("FFT3 round-trip [0,0,0]: got %v, want 1", real(v0)) } } func TestFFTN(t *testing.T) { a := mustFromFloats(t, []float64{ 1, 2, 3, 4, 5, 6, 7, 8, }, 2, 4) got, err := FFTN(a, nil) if err != nil { t.Fatal(err) } if got.Dtype() != core.Complex { t.Errorf("FFTN dtype: %s", got.Dtype()) } // Round-trip. back, err := IFFTN(got, nil) if err != nil { t.Fatal(err) } for i := range a.Len() { v, _ := core.ComplexAt(back, i/4, i%4) orig, _ := core.FloatAt(a, i/4, i%4) if math.Abs(real(v)-orig) > 1e-9 { t.Errorf("FFTN round-trip [%d]: got %v, want %v", i, real(v), orig) } } } func TestRFFT(t *testing.T) { // Real input of length 8 gives a spectrum of length 5. a := mustFromFloats(t, []float64{1, 0, 0, 0, 0, 0, 0, 0}, 8) spec, err := RFFT(a) if err != nil { t.Fatal(err) } if spec.Len() != 5 { t.Errorf("RFFT spectrum length: %d, want 5", spec.Len()) } // IRFFT round-trip. back, err := IRFFT(spec, 8) if err != nil { t.Fatal(err) } for i := range 8 { v, _ := core.FloatAt(back, i) orig, _ := core.FloatAt(a, i) if math.Abs(v-orig) > 1e-9 { t.Errorf("RFFT round-trip [%d]: got %v, want %v", i, v, orig) } } } func TestFFTFreq(t *testing.T) { f := FFTFreq(8, 1.0) if f.Len() != 8 { t.Errorf("FFTFreq length: %d", f.Len()) } // Frequencies: [0, 1/8, 2/8, 3/8, -4/8, -3/8, -2/8, -1/8] = [0, 0.125, 0.25, 0.375, -0.5, -0.375, -0.25, -0.125] want := []float64{0, 0.125, 0.25, 0.375, -0.5, -0.375, -0.25, -0.125} for i, w := range want { v, _ := core.FloatAt(f, i) if math.Abs(v-w) > 1e-9 { t.Errorf("FFTFreq [%d]: got %v, want %v", i, v, w) } } } // TestIRFFTLengthOneSpectrum pins the degenerate spectrum: a single // DC bin defaults to n = 1 instead of n = 0, an empty spectrum is an // error, and a non-positive explicit n is an error. func TestIRFFTLengthOneSpectrum(t *testing.T) { spec := mustComplexes(t, []complex128{complex(5, 0)}, 1) back, err := IRFFT(spec, 0) if err != nil { t.Fatalf("IRFFT default n: %v", err) } if back.Len() != 1 || math.Abs(back.FloatAt(0)-5) > 1e-12 { t.Fatalf("IRFFT of a DC-only spectrum = %v, want [5]", back) } back, err = IRFFT(spec, 1) if err != nil { t.Fatalf("IRFFT n=1: %v", err) } if back.Len() != 1 || math.Abs(back.FloatAt(0)-5) > 1e-12 { t.Fatalf("IRFFT n=1 = %v, want [5]", back) } if _, err := IRFFT(spec, 3); err == nil { t.Fatal("expected an error for a spectrum length that mismatches n") } if _, err := IRFFT(mustComplexes(t, nil, 0), 0); err == nil { t.Fatal("expected an error for an empty spectrum") } if _, err := IRFFT(spec, -2); err == nil { t.Fatal("expected an error for a negative n") } } // TestFFTEmptyComplexIsError pins the empty guard for complex input, // which bypasses the real-to-complex conversion where the check used // to live. func TestFFTEmptyComplexIsError(t *testing.T) { empty := mustComplexes(t, nil, 0) if _, err := FFT(empty); err == nil || !strings.Contains(err.Error(), "empty array") { t.Fatalf("FFT of an empty complex array: %v", err) } if _, err := IFFT(empty); err == nil || !strings.Contains(err.Error(), "empty array") { t.Fatalf("IFFT of an empty complex array: %v", err) } } // TestIFFTNSubsetDimsRoundTrip pins the inverse scaling: IFFTN divides // by the product of the transformed extents, so inverting a subset of // the dimensions restores the input exactly (scaling by the total // length would return a/2 here). func TestIFFTNSubsetDimsRoundTrip(t *testing.T) { src := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) fwd, err := FFTN(src, []int{1}) if err != nil { t.Fatalf("FFTN: %v", err) } back, err := IFFTN(fwd, []int{1}) if err != nil { t.Fatalf("IFFTN: %v", err) } for i := range 6 { want, _ := core.FloatAt(src, i) got, _ := core.ComplexAt(back, i) if math.Abs(real(got)-want) > 1e-9 || math.Abs(imag(got)) > 1e-9 { t.Fatalf("subset round-trip [%d]: got %v, want %v", i, got, want) } } // The full-dims path keeps its usual total scaling. fwdAll, err := FFTN(src, nil) if err != nil { t.Fatalf("FFTN all dims: %v", err) } backAll, err := IFFTN(fwdAll, nil) if err != nil { t.Fatalf("IFFTN all dims: %v", err) } for i := range 6 { want, _ := core.FloatAt(src, i) got, _ := core.ComplexAt(backAll, i) if math.Abs(real(got)-want) > 1e-9 || math.Abs(imag(got)) > 1e-9 { t.Fatalf("full round-trip [%d]: got %v, want %v", i, got, want) } } } // TestFFT2DoesNotMutateComplexInput pins the immutability contract for // the multi-dimension chain: FFT2 with an already complex input must // not overwrite the caller's payload. An in-place first pass used to // slip through here because the return values stayed correct; only the // caller's array silently held the spectrum afterwards. func TestFFT2DoesNotMutateComplexInput(t *testing.T) { vals := []complex128{ 1 + 2i, 3 - 1i, -0.5 + 0.5i, 2 + 0i, } a, err := core.FromComplexes(vals, 2, 2) if err != nil { t.Fatalf("FromComplexes: %v", err) } got, err := FFT2(a) if err != nil { t.Fatalf("FFT2: %v", err) } for i, want := range vals { if a.ComplexAt(i) != want { t.Fatalf("FFT2 mutated input[%d] = %v, want %v", i, a.ComplexAt(i), want) } } back, err := IFFT2(got) if err != nil { t.Fatalf("IFFT2: %v", err) } for i, want := range vals { if back.ComplexAt(i) != want { t.Fatalf("IFFT2 mutated gradient input or lost precision [%d] = %v, want %v", i, back.ComplexAt(i), want) } } }