33 lines
1.2 KiB
Python
33 lines
1.2 KiB
Python
from antlr4.error.ErrorListener import ErrorListener
|
|
import torch
|
|
|
|
def getIndexTensorAlongDim(tensor, dim):
|
|
shape = tensor.shape
|
|
|
|
# Create values: shape (size of dim)
|
|
values = torch.arange(shape[dim], dtype=torch.float32)
|
|
|
|
# Reshape values to align with the target dimension
|
|
view_shape = [1] * len(shape)
|
|
view_shape[dim] = shape[dim]
|
|
values = values.view(*view_shape)
|
|
|
|
# Broadcast to full shape
|
|
return values.expand(*shape)
|
|
|
|
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)
|
|
|
|
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):
|
|
raise ValueError(f"Syntax error in AudioExpr at line {line}, col {column}: {msg}")
|