From c24feebac163ca1aa58842ec01f61dc59c327fcf Mon Sep 17 00:00:00 2001 From: mcDandy Date: Wed, 3 Sep 2025 19:06:57 +0200 Subject: [PATCH] Fixed another bug with fft/ifft --- src/more_math/Parser/TensorEvalVisitor.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/src/more_math/Parser/TensorEvalVisitor.py b/src/more_math/Parser/TensorEvalVisitor.py index 94444dc..9a9b439 100644 --- a/src/more_math/Parser/TensorEvalVisitor.py +++ b/src/more_math/Parser/TensorEvalVisitor.py @@ -154,17 +154,20 @@ class TensorEvalVisitor(MathExprVisitor): 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['T'] = getIndexTensorAlongDim(shp_freq, 3) # now frame index + self.variables['S'] = getIndexTensorAlongDim(shp_freq, 3) # now frame index + self.variables['T'] = torch.full_like(shp_freq, shp_freq.shape[3]) # now num frames self.variables['B'] = getIndexTensorAlongDim(shp_freq, 0) self.variables['C'] = getIndexTensorAlongDim(shp_freq, 1) self.variables['F'] = getIndexTensorAlongDim(shp_freq, 2) + self.variables['K'] = shp_freq.shape[2] self.variables['R'] = torch.full_like(shp_freq, self.variables['R'].flatten()[0].item()) # Convert time→freq @@ -209,6 +212,7 @@ class TensorEvalVisitor(MathExprVisitor): self.variables['T'] = getIndexTensorAlongDim(shp_time, 2) self.variables['B'] = getIndexTensorAlongDim(shp_time, 0) self.variables['C'] = getIndexTensorAlongDim(shp_time, 1) + self.variables['S'] = getIndexTensorAlongDim(shp_freq, 2) self.variables['R'] = torch.full_like(shp_time, self.variables['R'].flatten()[0].item()) # Convert freq→time