Files
set-soft-AudioSeparation/source/inference/stft.py
T

165 lines
5.9 KiB
Python

# Short-Time Fourier Transform (STFT).
import logging
import numpy as np
import torch
from tqdm import tqdm
# Local imports
from ..utils.misc import NODES_NAME
from ..utils.torch import model_to_target
logger = logging.getLogger(f"{NODES_NAME}.stft")
class STFT:
def __init__(self, n_fft, hop_length, dim_f, device):
self.n_fft = n_fft
self.hop_length = hop_length
self.window = torch.hann_window(window_length=self.n_fft, periodic=True)
self.dim_f = dim_f
self.device = device
def __call__(self, x):
window = self.window.to(x.device)
batch_dims = x.shape[:-2]
c, t = x.shape[-2:]
x = x.reshape([-1, t])
x = torch.stft(x, n_fft=self.n_fft, hop_length=self.hop_length,
window=window, center=True, return_complex=True)
x = torch.view_as_real(x)
x = x.permute([0, 3, 1, 2])
x = x.reshape([*batch_dims, c, 2, -1, x.shape[-1]]
).reshape([*batch_dims, c * 2, -1, x.shape[-1]])
return x[..., :self.dim_f, :]
# Original code
# def inverse(self, x):
# window = self.window.to(x.device)
# batch_dims = x.shape[:-3]
# c, f, t = x.shape[-3:]
# n = self.n_fft // 2 + 1
# f_pad = torch.zeros([*batch_dims, c, n - f, t]).to(x.device)
# x = torch.cat([x, f_pad], -2)
# x = x.reshape([*batch_dims, c // 2, 2, n, t]).reshape([-1, 2, n, t])
# x = x.permute([0, 2, 3, 1])
# x = x[..., 0] + x[..., 1] * 1.j
# x = torch.istft(x, n_fft=self.n_fft,
# hop_length=self.hop_length, window=window, center=True)
# x = x.reshape([*batch_dims, 2, -1])
#
# return x
# Annotated code
def inverse(self, x):
"""
Correctly performs the inverse STFT.
x is the output of the model, shape (B, C, F, T)
With C == 4 (L/R as complex)
"""
window = self.window.to(x.device)
batch_dims = x.shape[:-3]
c, f, t = x.shape[-3:] # c is 4 here
assert c == 4
n = self.n_fft // 2 + 1 # Full number of frequency bins
# Pad the frequency dimension back to its original size
f_pad = torch.zeros([*batch_dims, c, n - f, t], device=x.device)
x = torch.cat([x, f_pad], -2)
# The key is to correctly un-stack the 4 channels back into (C, 2)
# where C=2 (stereo) and 2 is real/imag.
# Reshape (B, 4, F, T) -> (B, 2, 2, F, T)
# The new dimensions are (B, stereo_channels, real_imag, F, T)
x = x.reshape([*batch_dims, 2, 2, n, t])
# Permute to get (B, stereo_channels, F, T, real_imag)
x = x.permute(0, 1, 3, 4, 2)
# Ensure the tensor is contiguous in memory before the final view
x = x.contiguous()
# Now, view_as_complex will work on the last dimension
x = torch.view_as_complex(x) # Shape: (B, C, F, T) complex
# Reshape for istft: (B, C, F, T) -> (B*C, F, T)
x = x.reshape(-1, n, t)
# Perform inverse STFT
x = torch.istft(x, n_fft=self.n_fft, hop_length=self.hop_length, window=window, center=True)
# Reshape back to (B, C, num_samples)
x = x.reshape([*batch_dims, 2, -1])
return x
def stft_get_chunks(samples, n_fft, segment_size=256, hop_length=1024):
chunk_size = hop_length * (segment_size - 1)
step = chunk_size - n_fft
return 1 + (samples - chunk_size + step - 1) // step
def stft_chunk_process(waveform, d, model_run, device, segment_size=256, hop_length=1024, progress_bar_ui=None):
""" Inference helper for models using STFT information as input """
n_fft = d['mdx_n_fft_scale_set']
compensate = d['compensate']
mix_np = waveform.numpy()
stft = STFT(n_fft, hop_length, model_run.dim_f, device)
trim = n_fft // 2
chunk_size = hop_length * (segment_size - 1)
gen_size = chunk_size - 2 * trim
# --- Overlap-Add Loop ---
pad = gen_size + trim - (mix_np.shape[1] % gen_size)
# Padded mixture as a numpy array
# mixture = np.concatenate((np.zeros((2, trim)), mix_np, np.zeros((2, pad))), axis=1)
mixture = np.concatenate((np.zeros((2, trim), dtype=np.float32), mix_np, np.zeros((2, pad), dtype=np.float32)), axis=1)
step = chunk_size - n_fft # Correct step size for large overlap
result = np.zeros((1, 2, mixture.shape[1]), dtype=np.float32)
divider = np.zeros((1, 2, mixture.shape[1]), dtype=np.float32)
total_chunks = 1 + (mixture.shape[1] - chunk_size + step - 1) // step
logger.info(f"⚙️ Processing {total_chunks} chunks...")
model_run.target_device = device
with model_to_target(model_run):
for i in tqdm(range(0, mixture.shape[1] - chunk_size + 1, step)):
start = i
end = i + chunk_size
mix_part = mixture[:, start:end]
# Convert just the chunk to a tensor for the model
mix_part_tensor = torch.from_numpy(mix_part).unsqueeze(0).to(device)
spek = stft(mix_part_tensor)
spec_pred = model_run(spek)
# Get the output back as a numpy array
tar_waves_np = stft.inverse(spec_pred).cpu().detach().numpy()
# Hanning window applied to the output before adding
# window = np.hanning(chunk_size)
window = np.hanning(chunk_size).astype(np.float32)
window = np.tile(window[None, None, :], (1, 2, 1))
result[..., start:end] += tar_waves_np * window
divider[..., start:end] += window
if progress_bar_ui:
progress_bar_ui.update(1)
# --- Final Normalization and Trimming ---
divider[divider == 0] = 1.0
main_wav_np = (result[0] / divider[0]) # Get the 2D array
main_wav_np = main_wav_np[:, trim:-trim][:, :mix_np.shape[1]]
main_wav_np *= compensate
# Convert final result back to a torch tensor for saving
return torch.from_numpy(main_wav_np)