From e4faaded4b57a2eab80529a87d86cb43bd758b2e Mon Sep 17 00:00:00 2001 From: mcDandy Date: Fri, 5 Dec 2025 16:42:43 +0100 Subject: [PATCH] added fft/ifft compatibilty for other types --- README.md | 3 +- src/__init__.py | 0 src/more_math/ImageMathNode.py | 19 ++- src/more_math/LatentMathNode.py | 30 +++- src/more_math/Parser/TensorEvalVisitor.py | 182 ++++++++-------------- src/more_math/VideoMathNode.py | 20 ++- src/more_math/helper_functions.py | 80 ++-------- tests/conftest.py | 2 +- tests/test_more_math.py | 87 ++++++++--- 9 files changed, 196 insertions(+), 227 deletions(-) create mode 100644 src/__init__.py diff --git a/README.md b/README.md index d18047c..0da9b42 100644 --- a/README.md +++ b/README.md @@ -27,8 +27,7 @@ You can also get the node from comfy manager under the name of More math. - Trigonometric: `sin`, `cos`, `tan`, `asin`, `acos`, `atan`, `atan2` - Hyperbolic: `sinh`, `cosh`, `tanh`, `asinh`, `acosh`, `atanh` - Aggregates: `smin`, `smax` , `snorm` (scalar), `tmin`, `tmax`, `tnorm` (elementwise) -- Other: `floor`, `ceil`, `round`, `gamma`, `clamp`, `sigm` (sigmoid) -- Audio-specific: `fft` (short-time FFT), `ifft` (inverse sFFT **always return audio back to time domain before leaving node**), `angle` (in ifft only) +- Other: `floor`, `ceil`, `round`, `gamma`, `clamp`, `sigm` (sigmoid) `fft` (short-time FFT), `ifft` (inverse sFFT **always return data back to time or position domain before leaving node**), `angle` (in ifft only) ## Variables - **common inputs** (matches node input type): diff --git a/src/__init__.py b/src/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/more_math/ImageMathNode.py b/src/more_math/ImageMathNode.py index 6f5eedf..671e81c 100644 --- a/src/more_math/ImageMathNode.py +++ b/src/more_math/ImageMathNode.py @@ -58,11 +58,17 @@ class ImageMathNode(io.ComfyNode): b = torch.zeros_like(a) if b is None else b c = torch.zeros_like(a) if c is None else c d = torch.zeros_like(a) if d is None else d + + # permute to B, C, H, W + a = a.permute(0, 3, 1, 2) + b = b.permute(0, 3, 1, 2) + c = c.permute(0, 3, 1, 2) + d = d.permute(0, 3, 1, 2) B = getIndexTensorAlongDim(a, 0) - W = getIndexTensorAlongDim(a, 2) - H = getIndexTensorAlongDim(a, 1) - C = getIndexTensorAlongDim(a, 3) + C = getIndexTensorAlongDim(a, 1) + H = getIndexTensorAlongDim(a, 2) + W = getIndexTensorAlongDim(a, 3) variables = { 'a': a, 'b': b, 'c': c, 'd': d, @@ -70,10 +76,10 @@ class ImageMathNode(io.ComfyNode): 'X': W, 'Y': H, 'B': B,'batch': B, 'C': C,'channel': C, - 'W': a.shape[1], 'width': a.shape[1], + 'W': a.shape[3], 'width': a.shape[3], 'H': a.shape[2], 'height': a.shape[2], 'T': a.shape[0], 'batch_count': a.shape[0], - 'N': a.shape[3], 'channel_count': a.shape[3], + 'N': a.shape[1], 'channel_count': a.shape[1], } input_stream = InputStream(Image) lexer = MathExprLexer(input_stream) @@ -83,6 +89,9 @@ class ImageMathNode(io.ComfyNode): tree = parser.expr() visitor = TensorEvalVisitor(variables,a.shape) result = visitor.visit(tree) + + # permute back to B, H, W, C + result = result.permute(0, 2, 3, 1) return (result,) diff --git a/src/more_math/LatentMathNode.py b/src/more_math/LatentMathNode.py index 5095ae1..1dc244d 100644 --- a/src/more_math/LatentMathNode.py +++ b/src/more_math/LatentMathNode.py @@ -1,4 +1,5 @@ from inspect import cleandoc +from math import e from comfy_api.latest import ComfyExtension, io @@ -47,9 +48,9 @@ class LatentMathNode(io.ComfyNode): category="More math", inputs=[ io.Latent.Input(id="a"), - io.Latent.Input(id="b", optional=True), - io.Latent.Input(id="c", optional=True), - io.Latent.Input(id="d", optional=True), + io.Latent.Input(id="b", optional=True, lazy=True), + io.Latent.Input(id="c", optional=True, lazy=True), + io.Latent.Input(id="d", optional=True, lazy=True), io.Float.Input(id="w", default=0.0,optional=True, force_input=True), io.Float.Input(id="x", default=0.0,optional=True, force_input=True), io.Float.Input(id="y", default=0.0,optional=True, force_input=True), @@ -66,7 +67,30 @@ class LatentMathNode(io.ComfyNode): #OUTPUT_NODE = False #OUTPUT_TOOLTIPS = ("",) # Tooltips for the output node + async def check_lazy_status(self,Latent,a,b='',c='',d='',w='',x='',y='',z=''): + input_stream = InputStream(Latent) + lexer = MathExprLexer(input_stream) + stream = CommonTokenStream(lexer) + need_load = ['a'] if a is None else [] + for v in stream.getTokens(): + if v.type == MathExprParser.VARIABLE: + var_name = v.text + if var_name == 'b' and b is None: + need_load.append('b') + elif var_name == 'c' and c is None: + need_load.append('c') + elif var_name == 'd' and d is None: + need_load.append('d') + elif var_name == 'w' and w is None: + need_load.append('w') + elif var_name == 'x' and x is None: + need_load.append('x') + elif var_name == 'y' and y is None: + need_load.append('y') + elif var_name == 'z' and z is None: + need_load.append('z') + return need_load @classmethod def execute(cls, Latent, a, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0) -> io.NodeOutput: diff --git a/src/more_math/Parser/TensorEvalVisitor.py b/src/more_math/Parser/TensorEvalVisitor.py index d8e4d03..dba5ae1 100644 --- a/src/more_math/Parser/TensorEvalVisitor.py +++ b/src/more_math/Parser/TensorEvalVisitor.py @@ -9,6 +9,7 @@ from .MathExprVisitor import MathExprVisitor class TensorEvalVisitor(MathExprVisitor): def __init__(self, variables,shape): self.variables = variables + self.spatial_variables = variables.copy() self.shape = shape def visitNumberExp(self, ctx): @@ -130,124 +131,75 @@ class TensorEvalVisitor(MathExprVisitor): return val def visitSfftFunc(self, ctx): - hop_length = 256 - n_fft = 512 - - # Must start in freq-domain - if len(self.shape) != 4: - raise ValueError(f"SFFT input must be 4D freq-domain, got {len(self.shape)}D") - B, C, F, T_spec = self.shape - T = T_spec * hop_length - old_shape = self.shape - - # --- Switch to time-domain for child --- - time_shape = (B, C, T) - time_shape = self.variables['a'].shape if 'a' in self.variables else time_shape - self.shape = time_shape - shp_time = torch.zeros(time_shape) - self.variables['S'] = getIndexTensorAlongDim(shp_time, 2) - self.variables['sample'] = self.variables['S'] - self.variables['N'] = self.shape[1] - self.variables['channel_count'] = self.shape[1] - self.variables['T'] = torch.full_like(shp_time, self.shape[2]) - self.variables['sample_count'] = self.variables['T'] - self.variables['B'] = getIndexTensorAlongDim(shp_time, 0) - self.variables['batch'] = self.variables['B'] - self.variables['C'] = getIndexTensorAlongDim(shp_time, 1) - self.variables['channel'] = self.variables['C'] - self.variables['R'] = torch.full_like(shp_time, self.variables['R'].flatten()[0].item()) - self.variables['sample_rate'] = self.variables['R'] - - # Child sees time-domain - val = self.visit(ctx.expr()) - if val.ndim != 3: - raise ValueError(f"SFFT child must return 3D time-domain, got {val.ndim}D") - - # --- Back to freq-domain after child --- - num_freq_bins = n_fft // 2 + 1 - num_frames = T // hop_length + 1 - freq_shape = (B, C, num_freq_bins, num_frames) - freq_shape = old_shape - self.shape = freq_shape - shp_freq = torch.zeros(freq_shape, device=shp_time.device) - self.variables['S'] = getIndexTensorAlongDim(shp_freq, 3) # now frame index - self.variables['sample'] = self.variables['S'] - self.variables['T'] = torch.full_like(shp_freq, shp_freq.shape[3]) # now num frames - self.variables['sample_count'] = self.variables['T'] - self.variables['B'] = getIndexTensorAlongDim(shp_freq, 0) - self.variables['batch'] = self.variables['B'] - self.variables['C'] = getIndexTensorAlongDim(shp_freq, 1) - self.variables['channel'] = self.variables['C'] - self.variables['F'] = getIndexTensorAlongDim(shp_freq, 2) - self.variables['freqency'] = self.variables['F'] - self.variables['K'] = shp_freq.shape[2] - self.variables['frequency_count'] = self.variables['K'] - self.variables['R'] = torch.full_like(shp_freq, self.variables['R'].flatten()[0].item()) - self.variables['sample_rate'] = self.variables['R'] - self.variables['N'] = self.shape[1] - self.variables['channel_count'] = self.shape[1] - - # Convert time→freq - freq_val = time_to_freq(val, n_fft=n_fft, hop_length=hop_length) - - # Restore caller’s view - self.shape = old_shape - return freq_val + old_vars = self.variables + self.variables = self.spatial_variables.copy() + try: + val = self.visit(ctx.expr()) + return time_to_freq(val) + finally: + self.variables = old_vars def visitSifftFunc(self, ctx): - hop_length = 256 - n_fft = 512 - - # Must start in time-domain - if len(self.shape) != 3: - raise ValueError(f"SIFFT input must be 3D time-domain, got {len(self.shape)}D") - B, C, T = self.shape - old_shape = self.shape - - # --- Switch to freq-domain for child --- - num_freq_bins = n_fft // 2 + 1 - num_frames = T // hop_length + 1 - freq_shape = (B, C, num_freq_bins, num_frames) - self.shape = freq_shape - shp_freq = torch.zeros(freq_shape) - self.variables['K'] = getIndexTensorAlongDim(shp_freq, 2) - self.variables['frequency_count'] = self.variables['K'] - self.variables['F'] = shp_freq.shape[2] - self.variables['frequency'] = self.variables['F'] - self.variables['T'] = getIndexTensorAlongDim(shp_freq, 3) - self.variables['sample_count'] = self.variables['T'] - self.variables['B'] = getIndexTensorAlongDim(shp_freq, 0) - self.variables['batch'] = self.variables['B'] - self.variables['C'] = getIndexTensorAlongDim(shp_freq, 1) - self.variables['channel'] = self.variables['C'] - self.variables['S'] = getIndexTensorAlongDim(shp_freq, 2) - self.variables['sample'] = self.variables['S'] - self.variables['R'] = torch.full_like(shp_freq, self.variables['R'].flatten()[0].item()) - self.variables['sample_rate'] = self.variables['R'] - - # Child sees freq-domain - val = self.visit(ctx.expr()) - if val.ndim != 4: - raise ValueError(f"SIFFT child must return 4D freq-domain, got {val.ndim}D") - - # --- Back to time-domain after child --- - self.shape = old_shape - shp_time = torch.zeros(old_shape) - self.variables['T'] = getIndexTensorAlongDim(shp_time, 2) - self.variables['sample_count'] = self.variables['T'] - self.variables['B'] = getIndexTensorAlongDim(shp_time, 0) - self.variables['batch'] = self.variables['B'] - self.variables['C'] = getIndexTensorAlongDim(shp_time, 1) - self.variables['channel'] = self.variables['C'] - self.variables['S'] = getIndexTensorAlongDim(shp_freq, 2) - self.variables['sample'] = self.variables['S'] - self.variables['R'] = torch.full_like(shp_time, self.variables['R'].flatten()[0].item()) - self.variables['sample_rate'] = self.variables['R'] + old_vars = self.variables + # Switch to freq variables + self.variables = self.variables.copy() - # Convert freq→time - wav = freq_to_time(val, n_fft=n_fft, hop_length=hop_length, time=T) - - return wav + # Inject Freq variables based on shape + # Dimensions being transformed are 2 onwards + # We use a reference tensor from existing variables to get device/dtype if possible, + # or use val from a visit? We need vars BEFORE visit. + # We can construct index tensors using torch.arange like getIndexTensorAlongDim does. + # We need the device. 'a' is a safe bet for device source. + device = self.spatial_variables['a'].device if 'a' in self.spatial_variables else torch.device('cpu') + + dims = range(2, len(self.shape)) + for d in dims: + # Create index tensor for dim d + # Shape: ones with size at dim d + # getIndexTensorAlongDim logic: + # shape = tensor.shape + # values = torch.arange(shape[dim], ...) + # reshape and expand + size_d = self.shape[d] + values = torch.arange(size_d, dtype=torch.float32, device=device) + view_shape = [1] * len(self.shape) + view_shape[d] = size_d + values = values.view(*view_shape).expand(*self.shape) + + # Bind variables + if d == 2: + self.variables['K'] = values + self.variables['F'] = size_d + self.variables['Ky'] = values + self.variables['Fy'] = size_d + self.variables['frequency'] = self.variables['K'] # K is index + self.variables['frequency_count'] = self.variables['F'] # F is scalar + if d == 3: + self.variables['Kx'] = values + self.variables['Fx'] = size_d + + # Generic fallback + self.variables[f'K_dim{d}'] = values + self.variables[f'F_dim{d}'] = size_d + + # Calculate isotropic K (Euclidean distance from DC) + # K = sqrt(K_2^2 + K_3^2 + ...) + k_sq_sum = torch.zeros(self.shape, device=device) + dims = range(2, len(self.shape)) + for d in dims: + # Re-access the K variable for this dim (safe way) + k_val = self.variables.get(f'K_dim{d}') + if k_val is not None: + k_sq_sum = torch.add(k_sq_sum, torch.pow(k_val, 2)) + + self.variables['K'] = torch.sqrt(k_sq_sum) + self.variables['frequency'] = self.variables['K'] + + try: + val = self.visit(ctx.expr()) + return freq_to_time(val) + finally: + self.variables = old_vars # Two-argument functions def visitPowFunc(self, ctx): diff --git a/src/more_math/VideoMathNode.py b/src/more_math/VideoMathNode.py index 23b1b27..dcd61d1 100644 --- a/src/more_math/VideoMathNode.py +++ b/src/more_math/VideoMathNode.py @@ -73,12 +73,18 @@ class VideoMathNode(io.ComfyNode): dc = d.get_components() if d is not None else VideoComponents(images=torch.zeros_like(ac.images), audio={'waveform':torch.zeros_like(ac.audio['waveform']),'sample_rate':ac.audio['sample_rate']}, frame_rate=ac.frame_rate,metadata=None) + # permute images to B, C, H, W + ac.images = ac.images.permute(0, 3, 1, 2) + bc.images = bc.images.permute(0, 3, 1, 2) + cc.images = cc.images.permute(0, 3, 1, 2) + dc.images = dc.images.permute(0, 3, 1, 2) + B = getIndexTensorAlongDim(ac.images, 0) - X = getIndexTensorAlongDim(ac.images, 2) - Y = getIndexTensorAlongDim(ac.images, 1) - W = torch.full_like(Y, ac.images.shape[2], dtype=torch.float32) - H = torch.full_like(Y, ac.images.shape[1], dtype=torch.float32) - C = getIndexTensorAlongDim(ac.images, 3) + C = getIndexTensorAlongDim(ac.images, 1) + X = getIndexTensorAlongDim(ac.images, 3) # W + Y = getIndexTensorAlongDim(ac.images, 2) # H + W = torch.full_like(Y, ac.images.shape[3], dtype=torch.float32) + H = torch.full_like(Y, ac.images.shape[2], dtype=torch.float32) R = torch.full_like(Y, float(ac.frame_rate), dtype=torch.float32) T = torch.full_like(Y, ac.images.shape[0], dtype=torch.float32) @@ -90,7 +96,7 @@ class VideoMathNode(io.ComfyNode): 'C':C,'channel':C, 'R':R,'frame_rate':R, 'frame_count':ac.images.shape[0], - 'N':ac.images.shape[3],'channel_count':ac.images.shape[3]} + 'N':ac.images.shape[1],'channel_count':ac.images.shape[1]} input_stream = InputStream(Images) lexer = MathExprLexer(input_stream) @@ -100,6 +106,8 @@ class VideoMathNode(io.ComfyNode): tree = parser.expr() visitor = TensorEvalVisitor(variables,ac.images.shape) imgs = visitor.visit(tree) + # permute back to B, H, W, C + imgs = imgs.permute(0, 2, 3, 1) B = getIndexTensorAlongDim(ac.audio['waveform'], 0) diff --git a/src/more_math/helper_functions.py b/src/more_math/helper_functions.py index 051aa0f..cd34c0d 100644 --- a/src/more_math/helper_functions.py +++ b/src/more_math/helper_functions.py @@ -15,77 +15,17 @@ def getIndexTensorAlongDim(tensor, dim): # Broadcast to full shape return values.expand(*shape) -def time_to_freq(audio_dict: torch.Tensor, n_fft: int = 512, hop_length: int = 256) -> torch.Tensor: - waveform = audio_dict +def time_to_freq(element: torch.Tensor) -> torch.Tensor: + if element.ndim < 2: + raise ValueError("FFT requires at least 2 dimensions (Batch, Channel)") + dims = tuple(range(2, element.ndim)) + return torch.fft.fftn(element, dim=dims) - # Ensure the waveform is 3D with shape [batch, channels, time] - if waveform.ndimension() != 3: - raise ValueError(f"Expected 3D tensor, got {waveform.ndimension()}D tensor") - - # Create Hann window - window = torch.hann_window(n_fft, device=waveform.device) - - # Initialize the output spectrogram - spectrogram = torch.zeros( - waveform.shape[0], - waveform.shape[1], - n_fft // 2 + 1, # Frequency bins - waveform.shape[2]//hop_length+1, # Time steps - dtype=torch.complex64, - device=waveform.device - ) - - # Compute STFT on the last dimension for each value - for b in range(waveform.shape[0]): - for c in range(waveform.shape[1]): - stft_result = torch.stft( - waveform[b, c], - n_fft=n_fft, - hop_length=hop_length, - win_length=n_fft, - window=window, - center=True, - return_complex=True - ) - spectrogram[b, c] = stft_result - - return spectrogram - -def freq_to_time(freq_dict: torch.Tensor, n_fft: int = 512, hop_length: int = 256, time = None) -> torch.Tensor: - spectrogram = freq_dict - - # Ensure the spectrogram is 3D with shape [batch, channels, freq, time] - if spectrogram.ndimension() != 4: - raise ValueError(f"Expected 4D tensor, got {spectrogram.ndimension()}D tensor") - - # Initialize the output waveform - waveform = torch.zeros( - spectrogram.shape[0], - spectrogram.shape[1], - time if time else spectrogram.shape[3] * hop_length, # Total time samples - dtype=torch.float32, - device=spectrogram.device - ) - - # Create Hann window - window = torch.hann_window(n_fft, device=spectrogram.device) - - # Compute iSTFT on the last dimension for each value - for b in range(spectrogram.shape[0]): - for c in range(spectrogram.shape[1]): - istft_result = torch.istft( - spectrogram[b, c], - n_fft=n_fft, - hop_length=hop_length, - win_length=n_fft, - window=window, - center=True, - normalized=False, - length= time if time else spectrogram.shape[3] * hop_length - ) - waveform[b, c] = istft_result - - return waveform +def freq_to_time(element: torch.Tensor) -> torch.Tensor: + if element.ndim < 2: + raise ValueError("IFFT requires at least 2 dimensions (Batch, Channel)") + dims = tuple(range(2, element.ndim)) + return torch.fft.ifftn(element, dim=dims).real class ThrowingErrorListener(ErrorListener): def syntaxError(self, recognizer, offendingSymbol, line, column, msg, e): diff --git a/tests/conftest.py b/tests/conftest.py index 310609c..12af46a 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -3,4 +3,4 @@ import sys # Add the project root directory to Python path # This allows the tests to import the project -sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))) +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '../src'))) diff --git a/tests/test_more_math.py b/tests/test_more_math.py index aeadbcc..879fed9 100644 --- a/tests/test_more_math.py +++ b/tests/test_more_math.py @@ -2,10 +2,11 @@ """Tests for `more_math` package.""" -import pytest -from src.more_math.ConditioningMathNode import ConditioningMathNode -from src.more_math.LatentMathNode import LatentMathNode -from src.more_math.ImageMathNode import ImageMathNode +import unittest +import torch +from more_math.ConditioningMathNode import ConditioningMathNode +from more_math.LatentMathNode import LatentMathNode +from more_math.ImageMathNode import ImageMathNode import tokenize from io import StringIO @@ -25,31 +26,67 @@ def tokenize_expression(expr): filtered_tokens.append((token_name, tokval.strip())) return filtered_tokens -def test_conditioning_math_node_initialization(): - node = ConditioningMathNode() - assert isinstance(node, ConditioningMathNode) -def test_conditioning_math_node_metadata(): - assert ConditioningMathNode.RETURN_TYPES == ("CONDITIONING",) - assert ConditioningMathNode.FUNCTION == "condMathNode" - assert ConditioningMathNode.CATEGORY == "More math" +class TestMoreMath(unittest.TestCase): + def test_conditioning_math_node_initialization(self): + node = ConditioningMathNode() + self.assertIsInstance(node, ConditioningMathNode) -def test_latent_math_node_initialization(): - node = LatentMathNode() - assert isinstance(node, LatentMathNode) + def test_conditioning_math_node_metadata(self): + self.assertEqual(ConditioningMathNode.RETURN_TYPES, ["CONDITIONING"]) + self.assertEqual(ConditioningMathNode.FUNCTION, "EXECUTE_NORMALIZED") + self.assertEqual(ConditioningMathNode.CATEGORY, "More math") -def test_latent_math_node_metadata(): - assert LatentMathNode.RETURN_TYPES == ("LATENT",) - assert LatentMathNode.FUNCTION == "latMathNode" - assert LatentMathNode.CATEGORY == "More math" + def test_latent_math_node_initialization(self): + node = LatentMathNode() + self.assertIsInstance(node, LatentMathNode) -def test_image_math_node_initialization(): - node = ImageMathNode() - assert isinstance(node, ImageMathNode) + def test_latent_math_node_metadata(self): + self.assertEqual(LatentMathNode.RETURN_TYPES, ["LATENT"]) + self.assertEqual(LatentMathNode.FUNCTION, "EXECUTE_NORMALIZED") + self.assertEqual(LatentMathNode.CATEGORY, "More math") -def test_image_math_node_metadata(): - assert ImageMathNode.RETURN_TYPES == ("IMAGE",) - assert ImageMathNode.FUNCTION == "imgMathNode" - assert ImageMathNode.CATEGORY == "More math" + def test_image_math_node_initialization(self): + node = ImageMathNode() + self.assertIsInstance(node, ImageMathNode) + + def test_image_math_node_metadata(self): + self.assertEqual(ImageMathNode.RETURN_TYPES, ["IMAGE"]) + self.assertEqual(ImageMathNode.FUNCTION, "EXECUTE_NORMALIZED") + self.assertEqual(ImageMathNode.CATEGORY, "More math") + + def test_fft_invertibility(self): + # 1. Create random input latent (Batch, Channel, Height, Width) + input_tensor = torch.randn(1, 4, 32, 32, dtype=torch.float32) + input_dict = {"samples": input_tensor} + + # 2. Execute ifft(fft(a)) + # Note: execute is a classmethod + result = LatentMathNode.execute( + Latent="ifft(fft(a))", + a=input_dict + ) + + # 3. Get output tensor + output_tensor = result[0]["samples"] + + # 4. Check correctness (approximate equality) + self.assertTrue(torch.allclose(input_tensor, output_tensor, atol=1e-5), \ + f"Max difference: {(input_tensor - output_tensor).abs().max()}") + + def test_image_fft_dims(self): + # Image input is (Batch, Height, Width, Channel) + # Verify 2D FFT works by doing a round trip + input_tensor = torch.randn(1, 32, 32, 3, dtype=torch.float32) + + result = ImageMathNode.execute( + Image="ifft(fft(a))", + a=input_tensor + ) + output_tensor = result[0] + + self.assertEqual(input_tensor.shape, output_tensor.shape) + self.assertTrue(torch.allclose(input_tensor, output_tensor, atol=1e-5), \ + f"Image FFT round trip failed. Max diff: {(input_tensor - output_tensor).abs().max()}")