Files
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

248 lines
6.6 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package signal
import "sourcedock.dev/petrbalvin/tensor/internal/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)
}
}
}