mirror of
https://github.com/RVC-Project/Retrieval-based-Voice-Conversion-WebUI.git
synced 2026-08-29 10:09:32 +02:00
1110 lines
36 KiB
Python
1110 lines
36 KiB
Python
from contextlib import contextmanager, nullcontext
|
|
|
|
import numpy as np
|
|
import torch
|
|
import torch.nn as nn
|
|
from numpy.typing import NDArray
|
|
from typing import Dict
|
|
|
|
from pymss_core import get_model_from_config as _core_get_model_from_config
|
|
|
|
from .config import load_config
|
|
from .progress import _ProgressContext
|
|
|
|
|
|
def _model_target(model):
|
|
return model.module if isinstance(model, nn.DataParallel) else model
|
|
|
|
|
|
def get_model_from_config(model_type, config_path, model_kwargs_override=None):
|
|
"""Instantiate a separation model from a loaded model config.
|
|
|
|
Args:
|
|
model_type (Any): Model type value.
|
|
config_path (str | os.PathLike | None): Config path value.
|
|
model_kwargs_override (Any, optional): Model kwargs override value. Defaults to None.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
if model_type == "mel_band_roformer":
|
|
model_kwargs_override = dict(model_kwargs_override or {})
|
|
model_kwargs_override.setdefault("zero_dc", False)
|
|
return _core_get_model_from_config(
|
|
model_type, config_path, model_kwargs_override=model_kwargs_override
|
|
)
|
|
if model_type == "bandit_v2":
|
|
config = load_config(config_path)
|
|
from .modules.bandit_v2.bandit import Bandit
|
|
|
|
return Bandit(**config.kwargs), config
|
|
return _core_get_model_from_config(model_type, config_path, model_kwargs_override=model_kwargs_override)
|
|
|
|
|
|
def clear_mlx_cache():
|
|
"""Clear MLX memory caches when the MLX backend is available.
|
|
|
|
Args:
|
|
None: This callable does not accept user-provided arguments.
|
|
|
|
Returns:
|
|
None: This callable completes for its side effects."""
|
|
try:
|
|
import mlx.core as mx
|
|
except Exception:
|
|
return
|
|
|
|
clear_cache = getattr(mx, "clear_cache", None)
|
|
if clear_cache is None:
|
|
clear_cache = getattr(getattr(mx, "metal", None), "clear_cache", None)
|
|
if clear_cache is not None:
|
|
clear_cache()
|
|
|
|
|
|
def _getWindowingArray(window_size, fade_size):
|
|
"""Implement the getWindowingArray helper.
|
|
|
|
Args:
|
|
window_size (Any): Window size value.
|
|
fade_size (Any): Fade size value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
if fade_size <= 0:
|
|
return torch.ones(window_size)
|
|
|
|
fadein = torch.linspace(0, 1, fade_size)
|
|
fadeout = torch.linspace(1, 0, fade_size)
|
|
window = torch.ones(window_size)
|
|
window[-fade_size:] *= fadeout
|
|
window[:fade_size] *= fadein
|
|
return window
|
|
|
|
|
|
def _build_chunk_plan(total_length, chunk_size, step, fade_size):
|
|
"""Build chunk plan.
|
|
|
|
Args:
|
|
total_length (Any): Total length value.
|
|
chunk_size (Any): Chunk size value.
|
|
step (Any): Step value.
|
|
fade_size (Any): Fade size value.
|
|
|
|
Returns:
|
|
Any: Built value."""
|
|
starts = list(range(0, total_length, step))
|
|
normal_window = _getWindowingArray(chunk_size, fade_size)
|
|
|
|
def window_for(start):
|
|
"""Implement the window for helper.
|
|
|
|
Args:
|
|
start (Any): Start value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
length = min(chunk_size, total_length - start)
|
|
if start != 0 and start + length < total_length:
|
|
return normal_window
|
|
window = normal_window.clone()
|
|
if start == 0:
|
|
window[:fade_size] = 1
|
|
if start + length >= total_length:
|
|
window[max(0, length - fade_size) : length] = 1
|
|
return window
|
|
|
|
return starts, [window_for(start) for start in starts]
|
|
|
|
|
|
def _get_inference_step(config, chunk_size):
|
|
"""Return inference step.
|
|
|
|
Args:
|
|
config (AttrDict | dict): Loaded pymss configuration.
|
|
chunk_size (Any): Chunk size value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
overlap_size = int(config.inference.get("overlap_size", chunk_size // 2))
|
|
if overlap_size < 0 or overlap_size >= chunk_size:
|
|
raise ValueError("inference.overlap_size must be >= 0 and < audio.chunk_size")
|
|
return chunk_size - overlap_size
|
|
|
|
|
|
def _complete_chunk_count(total_length, chunk_size, step):
|
|
"""Implement the complete chunk count helper.
|
|
|
|
Args:
|
|
total_length (Any): Total length value.
|
|
chunk_size (Any): Chunk size value.
|
|
step (Any): Step value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
return 0 if total_length < chunk_size else (total_length - chunk_size) // step + 1
|
|
|
|
|
|
def _fold_windows(counter, windows, step, start_offset=0):
|
|
"""Implement the fold windows helper.
|
|
|
|
Args:
|
|
counter (Any): Counter value.
|
|
windows (Any): Windows value.
|
|
step (Any): Step value.
|
|
start_offset (Any, optional): Start offset value. Defaults to 0.
|
|
|
|
Returns:
|
|
None: This callable completes for its side effects."""
|
|
n_chunks = windows.shape[0]
|
|
if n_chunks == 0:
|
|
return
|
|
|
|
chunk_size = windows.shape[-1]
|
|
output_length = (n_chunks - 1) * step + chunk_size
|
|
folded_counter = nn.functional.fold(
|
|
windows.transpose(0, 1).unsqueeze(0),
|
|
output_size=(1, output_length),
|
|
kernel_size=(1, chunk_size),
|
|
stride=(1, step),
|
|
)
|
|
counter[..., start_offset : start_offset + output_length] += folded_counter.view(1, 1, output_length)
|
|
|
|
|
|
def _fold_chunk_batch(result, chunks, windows, step, start_offset=0):
|
|
"""Implement the fold chunk batch helper.
|
|
|
|
Args:
|
|
result (Any): Result value.
|
|
chunks (Any): Chunks value.
|
|
windows (Any): Windows value.
|
|
step (Any): Step value.
|
|
start_offset (Any, optional): Start offset value. Defaults to 0.
|
|
|
|
Returns:
|
|
None: This callable completes for its side effects."""
|
|
n_chunks = chunks.shape[0]
|
|
if n_chunks == 0:
|
|
return
|
|
|
|
chunk_size = chunks.shape[-1]
|
|
output_length = (n_chunks - 1) * step + chunk_size
|
|
n_sources, n_channels = chunks.shape[1:3]
|
|
|
|
folded = nn.functional.fold(
|
|
(chunks * windows[:, None, None, :]).permute(1, 2, 3, 0).reshape(1, n_sources * n_channels * chunk_size, n_chunks),
|
|
output_size=(1, output_length),
|
|
kernel_size=(1, chunk_size),
|
|
stride=(1, step),
|
|
)
|
|
result[..., start_offset : start_offset + output_length] += folded.view(n_sources, n_channels, output_length)
|
|
|
|
|
|
def _ensure_source_dim(x, chunk_batch):
|
|
"""Ensure source dim.
|
|
|
|
Args:
|
|
x (Any): X value.
|
|
chunk_batch (Any): Chunk batch value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
return x.unsqueeze(1) if x.ndim == chunk_batch.ndim else x
|
|
|
|
|
|
def _fit_tensor_length(x, length):
|
|
"""Implement the fit tensor length helper.
|
|
|
|
Args:
|
|
x (Any): X value.
|
|
length (Any): Length value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
if x.shape[-1] > length:
|
|
return x[..., :length]
|
|
if x.shape[-1] < length:
|
|
return nn.functional.pad(x, (0, length - x.shape[-1]))
|
|
return x
|
|
|
|
|
|
def _autocast(device, enabled):
|
|
"""Implement the autocast helper.
|
|
|
|
Args:
|
|
device (Any): Device value.
|
|
enabled (Any): Enabled value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
device_type = torch.device(device).type
|
|
if enabled and device_type in ("cuda", "mps"):
|
|
return torch.amp.autocast(device_type, dtype=torch.float16)
|
|
return nullcontext()
|
|
|
|
|
|
def _inference_context(device):
|
|
if torch.device(device).type == "privateuseone":
|
|
return torch.no_grad()
|
|
return torch.inference_mode()
|
|
|
|
|
|
def _source_names(config):
|
|
"""Implement the source names helper.
|
|
|
|
Args:
|
|
config (AttrDict | dict): Loaded pymss configuration.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
return config.training.instruments if config.training.target_instrument is None else [config.training.target_instrument]
|
|
|
|
|
|
def _normalize_source_indices(config, source_indices):
|
|
"""Normalize source indices.
|
|
|
|
Args:
|
|
config (AttrDict | dict): Loaded pymss configuration.
|
|
source_indices (Any): Source indices value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
if source_indices is None:
|
|
return None
|
|
source_count = len(_source_names(config))
|
|
indices = tuple(int(index) for index in source_indices)
|
|
if not indices:
|
|
raise ValueError("source_indices must not be empty")
|
|
if len(set(indices)) != len(indices):
|
|
raise ValueError("source_indices must not contain duplicates")
|
|
if min(indices) < 0 or max(indices) >= source_count:
|
|
raise ValueError(f"source_indices must be in range [0, {source_count})")
|
|
return indices
|
|
|
|
|
|
def _source_count(config, source_indices=None):
|
|
"""Implement the source count helper.
|
|
|
|
Args:
|
|
config (AttrDict | dict): Loaded pymss configuration.
|
|
source_indices (Any, optional): Source indices value. Defaults to None.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
return len(_source_names(config)) if source_indices is None else len(source_indices)
|
|
|
|
|
|
def _sources_to_dict(config, estimated_sources, source_indices=None):
|
|
"""Implement the sources to dict helper.
|
|
|
|
Args:
|
|
config (AttrDict | dict): Loaded pymss configuration.
|
|
estimated_sources (Any): Estimated sources value.
|
|
source_indices (Any, optional): Source indices value. Defaults to None.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
names = _source_names(config)
|
|
if source_indices is not None:
|
|
names = [names[index] for index in source_indices]
|
|
return {k: v for k, v in zip(names, estimated_sources)}
|
|
|
|
|
|
def _prepare_mix_for_chunks(mix, border):
|
|
"""Implement the prepare mix for chunks helper.
|
|
|
|
Args:
|
|
mix (np.ndarray): Mix value.
|
|
border (Any): Border value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
length_init = mix.shape[-1]
|
|
mix = mix.unsqueeze(0) if mix.ndim == 1 else mix
|
|
if length_init > 2 * border and border > 0:
|
|
mix = nn.functional.pad(mix, (border, border), mode="reflect")
|
|
return mix, length_init
|
|
|
|
|
|
def _init_overlap_buffers(config, mix, device, use_fast_path, source_indices=None):
|
|
"""Implement the init overlap buffers helper.
|
|
|
|
Args:
|
|
config (AttrDict | dict): Loaded pymss configuration.
|
|
mix (np.ndarray): Mix value.
|
|
device (Any): Device value.
|
|
use_fast_path (Any): Use fast path value.
|
|
source_indices (Any, optional): Source indices value. Defaults to None.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
req_shape = (_source_count(config, source_indices),) + tuple(mix.shape)
|
|
result_device = device if use_fast_path else "cpu"
|
|
counter_shape = (1, 1, mix.shape[1])
|
|
result = torch.zeros(req_shape, dtype=torch.float32, device=result_device)
|
|
counter = torch.zeros(counter_shape, dtype=torch.float32, device=result_device)
|
|
return result, counter
|
|
|
|
|
|
def _model_mix(mix, device):
|
|
"""Implement the model mix helper.
|
|
|
|
Args:
|
|
mix (np.ndarray): Mix value.
|
|
device (Any): Device value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
return mix.to(device) if torch.device(device).type != "cpu" else mix
|
|
|
|
|
|
@contextmanager
|
|
def _model_source_context(model, source_indices):
|
|
"""Implement the model source context helper.
|
|
|
|
Args:
|
|
model (str): Model value.
|
|
source_indices (Any): Source indices value.
|
|
|
|
Returns:
|
|
None: This callable completes for its side effects."""
|
|
target = _model_target(model)
|
|
sentinel = object()
|
|
previous = getattr(target, "_pymss_source_indices", sentinel)
|
|
if source_indices is not None:
|
|
target._pymss_source_indices = source_indices
|
|
try:
|
|
yield
|
|
finally:
|
|
if previous is sentinel:
|
|
if hasattr(target, "_pymss_source_indices"):
|
|
delattr(target, "_pymss_source_indices")
|
|
else:
|
|
target._pymss_source_indices = previous
|
|
|
|
|
|
def _select_sources(chunks, source_indices, already_selected=False):
|
|
"""Implement the select sources helper.
|
|
|
|
Args:
|
|
chunks (Any): Chunks value.
|
|
source_indices (Any): Source indices value.
|
|
already_selected (Any, optional): Already selected value. Defaults to False.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
if source_indices is None or already_selected:
|
|
return chunks
|
|
index = torch.as_tensor(source_indices, device=chunks.device)
|
|
return chunks.index_select(1, index)
|
|
|
|
|
|
def _run_model_chunk(model, arr, chunk_size, source_indices=None):
|
|
"""Run model chunk.
|
|
|
|
Args:
|
|
model (str): Model value.
|
|
arr (np.ndarray): Arr value.
|
|
chunk_size (Any): Chunk size value.
|
|
source_indices (Any, optional): Source indices value. Defaults to None.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
target = _model_target(model)
|
|
chunks = _fit_tensor_length(_ensure_source_dim(model(arr), arr).float(), chunk_size)
|
|
already_selected = (
|
|
source_indices is not None and hasattr(target, "_active_source_indices") and chunks.shape[1] == len(source_indices)
|
|
)
|
|
return _select_sources(chunks, source_indices, already_selected=already_selected)
|
|
|
|
|
|
def _extract_chunk(mix, start, chunk_size):
|
|
"""Implement the extract chunk helper.
|
|
|
|
Args:
|
|
mix (np.ndarray): Mix value.
|
|
start (Any): Start value.
|
|
chunk_size (Any): Chunk size value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
length = min(chunk_size, mix.shape[1] - start)
|
|
part = mix[:, start : start + chunk_size]
|
|
if length == chunk_size:
|
|
return part, length
|
|
if length > chunk_size // 2 + 1:
|
|
part = nn.functional.pad(part, (0, chunk_size - length), mode="reflect")
|
|
else:
|
|
part = nn.functional.pad(part, (0, chunk_size - length, 0, 0), mode="constant", value=0)
|
|
return part, length
|
|
|
|
|
|
def _add_weighted_chunk(result, counter, chunk, window, start, length):
|
|
"""Implement the add weighted chunk helper.
|
|
|
|
Args:
|
|
result (Any): Result value.
|
|
counter (Any): Counter value.
|
|
chunk (Any): Chunk value.
|
|
window (Any): Window value.
|
|
start (Any): Start value.
|
|
length (Any): Length value.
|
|
|
|
Returns:
|
|
None: This callable completes for its side effects."""
|
|
device = result.device
|
|
window = window.to(device=device, dtype=torch.float32)[:length]
|
|
result[..., start : start + length] += chunk[..., :length].to(device=device, dtype=torch.float32) * window
|
|
counter[..., start : start + length] += window
|
|
|
|
|
|
def _run_complete_chunks(
|
|
model,
|
|
mix,
|
|
windows,
|
|
result,
|
|
counter,
|
|
chunk_size,
|
|
step,
|
|
batch_size,
|
|
progress,
|
|
source_indices=None,
|
|
):
|
|
"""Run complete chunks.
|
|
|
|
Args:
|
|
model (str): Model value.
|
|
mix (np.ndarray): Mix value.
|
|
windows (Any): Windows value.
|
|
result (Any): Result value.
|
|
counter (Any): Counter value.
|
|
chunk_size (Any): Chunk size value.
|
|
step (Any): Step value.
|
|
batch_size (Any): Batch size value.
|
|
progress (Any): Progress value.
|
|
source_indices (Any, optional): Source indices value. Defaults to None.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
n_chunks = _complete_chunk_count(mix.shape[1], chunk_size, step)
|
|
if n_chunks == 0:
|
|
return 0
|
|
|
|
n_complete = n_chunks
|
|
if len(windows) > n_chunks:
|
|
n_complete -= n_complete % batch_size
|
|
if n_complete == 0:
|
|
return 0
|
|
|
|
inputs = mix.unfold(-1, chunk_size, step).permute(1, 0, 2)[:n_complete]
|
|
fold_windows = torch.stack(windows[:n_complete], dim=0).to(device=result.device, dtype=torch.float32)
|
|
_fold_windows(counter, fold_windows, step)
|
|
|
|
for batch_start in range(0, n_complete, batch_size):
|
|
batch_end = min(batch_start + batch_size, n_complete)
|
|
chunks = _run_model_chunk(model, inputs[batch_start:batch_end].contiguous(), chunk_size, source_indices)
|
|
_fold_chunk_batch(
|
|
result,
|
|
chunks,
|
|
fold_windows[batch_start:batch_end],
|
|
step,
|
|
start_offset=batch_start * step,
|
|
)
|
|
progress.update(step * (batch_end - batch_start))
|
|
|
|
return n_complete
|
|
|
|
|
|
def _run_tail_chunks(
|
|
model,
|
|
mix,
|
|
starts,
|
|
windows,
|
|
result,
|
|
counter,
|
|
chunk_size,
|
|
step,
|
|
batch_size,
|
|
first_chunk,
|
|
progress,
|
|
source_indices=None,
|
|
):
|
|
"""Run tail chunks.
|
|
|
|
Args:
|
|
model (str): Model value.
|
|
mix (np.ndarray): Mix value.
|
|
starts (Any): Starts value.
|
|
windows (Any): Windows value.
|
|
result (Any): Result value.
|
|
counter (Any): Counter value.
|
|
chunk_size (Any): Chunk size value.
|
|
step (Any): Step value.
|
|
batch_size (Any): Batch size value.
|
|
first_chunk (Any): First chunk value.
|
|
progress (Any): Progress value.
|
|
source_indices (Any, optional): Source indices value. Defaults to None.
|
|
|
|
Returns:
|
|
None: This callable completes for its side effects."""
|
|
for batch_start in range(first_chunk, len(starts), batch_size):
|
|
batch_indices = range(batch_start, min(batch_start + batch_size, len(starts)))
|
|
batch = [(_extract_chunk(mix, starts[idx], chunk_size), idx) for idx in batch_indices]
|
|
batch_data = [chunk for (chunk, _), _ in batch]
|
|
chunks = _run_model_chunk(model, torch.stack(batch_data, dim=0), chunk_size, source_indices)
|
|
for j, ((_, length), idx) in enumerate(batch):
|
|
start = starts[idx]
|
|
_add_weighted_chunk(result, counter, chunks[j], windows[idx], start, length)
|
|
|
|
progress.update(step * len(batch_data))
|
|
|
|
|
|
def _finalize_overlap(result, counter, length_init, border):
|
|
"""Implement the finalize overlap helper.
|
|
|
|
Args:
|
|
result (Any): Result value.
|
|
counter (Any): Counter value.
|
|
length_init (Any): Length init value.
|
|
border (Any): Border value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
if length_init > 2 * border and border > 0:
|
|
start, end = border, border + length_init
|
|
else:
|
|
start, end = 0, result.shape[-1]
|
|
|
|
result = result[..., start:end]
|
|
counter = counter[..., start:end]
|
|
output_shape = result.shape[:-1] + (end - start,)
|
|
|
|
if torch.device(result.device).type != "cuda":
|
|
estimated_sources = (result / counter).cpu().numpy()
|
|
np.nan_to_num(estimated_sources, copy=False, nan=0.0)
|
|
return estimated_sources
|
|
|
|
counter_min, counter_max = torch.aminmax(counter)
|
|
divide_counter = bool((counter_min - 1).abs().item() > 1e-6 or (counter_max - 1).abs().item() > 1e-6)
|
|
samples_per_chunk = max(1, (512 * 1024 * 1024) // (max(1, result.shape[0] * result.shape[1]) * 4))
|
|
estimated_sources_t = torch.empty(output_shape, dtype=torch.float32, device="cpu")
|
|
for offset in range(0, result.shape[-1], samples_per_chunk):
|
|
chunk_end = min(offset + samples_per_chunk, result.shape[-1])
|
|
source = result[..., offset:chunk_end]
|
|
if divide_counter:
|
|
source = source / counter[..., offset:chunk_end]
|
|
estimated_sources_t[..., offset:chunk_end].copy_(source)
|
|
estimated_sources = estimated_sources_t.numpy()
|
|
if divide_counter:
|
|
np.nan_to_num(estimated_sources, copy=False, nan=0.0)
|
|
return estimated_sources
|
|
|
|
|
|
def _mlx_reflect_pad_1d(x, left=0, right=0):
|
|
"""Implement the mlx reflect pad 1d helper.
|
|
|
|
Args:
|
|
x (Any): X value.
|
|
left (Any, optional): Left value. Defaults to 0.
|
|
right (Any, optional): Right value. Defaults to 0.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
import mlx.core as mx
|
|
|
|
parts = []
|
|
if left > 0:
|
|
parts.append(x[..., 1 : left + 1][..., ::-1])
|
|
parts.append(x)
|
|
if right > 0:
|
|
parts.append(x[..., -right - 1 : -1][..., ::-1])
|
|
return mx.concatenate(parts, axis=-1)
|
|
|
|
|
|
def _mlx_get_windowing_array(window_size, fade_size):
|
|
"""Implement the mlx get windowing array helper.
|
|
|
|
Args:
|
|
window_size (Any): Window size value.
|
|
fade_size (Any): Fade size value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
import mlx.core as mx
|
|
|
|
if fade_size <= 0:
|
|
return mx.ones((window_size,), dtype=mx.float32)
|
|
fadein = mx.linspace(0, 1, fade_size)
|
|
fadeout = mx.linspace(1, 0, fade_size)
|
|
window = mx.ones((window_size,), dtype=mx.float32)
|
|
window = window.at[:fade_size].multiply(fadein)
|
|
window = window.at[-fade_size:].multiply(fadeout)
|
|
return window
|
|
|
|
|
|
def _mlx_build_chunk_plan(total_length, chunk_size, step, fade_size):
|
|
"""Implement the mlx build chunk plan helper.
|
|
|
|
Args:
|
|
total_length (Any): Total length value.
|
|
chunk_size (Any): Chunk size value.
|
|
step (Any): Step value.
|
|
fade_size (Any): Fade size value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
starts = list(range(0, total_length, step))
|
|
normal_window = _mlx_get_windowing_array(chunk_size, fade_size)
|
|
windows = []
|
|
for start in starts:
|
|
length = min(chunk_size, total_length - start)
|
|
if start != 0 and start + length < total_length:
|
|
windows.append(normal_window)
|
|
continue
|
|
window = normal_window
|
|
if start == 0 and fade_size > 0:
|
|
window = window.at[:fade_size].add(1 - window[:fade_size])
|
|
if start + length >= total_length and fade_size > 0:
|
|
tail = slice(max(0, length - fade_size), length)
|
|
window = window.at[tail].add(1 - window[tail])
|
|
windows.append(window)
|
|
return starts, windows
|
|
|
|
|
|
def _mlx_prepare_mix_for_chunks(mix, border):
|
|
"""Implement the mlx prepare mix for chunks helper.
|
|
|
|
Args:
|
|
mix (np.ndarray): Mix value.
|
|
border (Any): Border value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
import mlx.core as mx
|
|
|
|
length_init = mix.shape[-1]
|
|
mix = mx.array(np.asarray(mix, dtype=np.float32))
|
|
if mix.ndim == 1:
|
|
mix = mix[None, :]
|
|
if length_init > 2 * border and border > 0:
|
|
mix = _mlx_reflect_pad_1d(mix, border, border)
|
|
return mix, length_init
|
|
|
|
|
|
def _mlx_extract_chunk(mix, start, chunk_size):
|
|
"""Implement the mlx extract chunk helper.
|
|
|
|
Args:
|
|
mix (np.ndarray): Mix value.
|
|
start (Any): Start value.
|
|
chunk_size (Any): Chunk size value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
import mlx.core as mx
|
|
|
|
length = min(chunk_size, mix.shape[1] - start)
|
|
part = mix[:, start : start + chunk_size]
|
|
if length == chunk_size:
|
|
return part, length
|
|
pad = chunk_size - length
|
|
if length > chunk_size // 2 + 1:
|
|
part = _mlx_reflect_pad_1d(part, right=pad)
|
|
else:
|
|
part = mx.pad(part, [(0, 0), (0, pad)])
|
|
return part, length
|
|
|
|
|
|
def _mlx_fit_length(x, length):
|
|
"""Implement the mlx fit length helper.
|
|
|
|
Args:
|
|
x (Any): X value.
|
|
length (Any): Length value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
import mlx.core as mx
|
|
|
|
if x.shape[-1] > length:
|
|
return x[..., :length]
|
|
if x.shape[-1] < length:
|
|
return mx.pad(x, [(0, 0)] * (x.ndim - 1) + [(0, length - x.shape[-1])])
|
|
return x
|
|
|
|
|
|
@contextmanager
|
|
def _mlx_clear_cache_after_eval(enabled=False):
|
|
"""Clear MLX allocator cache after explicit eval points when requested."""
|
|
if not enabled:
|
|
yield
|
|
return
|
|
import mlx.core as mx
|
|
|
|
original_eval = mx.eval
|
|
|
|
def eval_and_clear(*args, **kwargs):
|
|
result = original_eval(*args, **kwargs)
|
|
clear_mlx_cache()
|
|
return result
|
|
|
|
mx.eval = eval_and_clear
|
|
try:
|
|
yield
|
|
finally:
|
|
mx.eval = original_eval
|
|
|
|
|
|
def _mlx_run_model_chunk(model, arr, chunk_size, clear_cache_after_eval=False):
|
|
"""Implement the mlx run model chunk helper.
|
|
|
|
Args:
|
|
model (str): Model value.
|
|
arr (np.ndarray): Arr value.
|
|
chunk_size (Any): Chunk size value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
with _mlx_clear_cache_after_eval(clear_cache_after_eval):
|
|
y = model.mlx_forward_mx(arr)
|
|
if y.ndim == arr.ndim:
|
|
y = y[:, None]
|
|
return _mlx_fit_length(y, chunk_size)
|
|
|
|
|
|
def _mlx_select_sources(chunks, source_indices):
|
|
"""Implement the mlx select sources helper.
|
|
|
|
Args:
|
|
chunks (Any): Chunks value.
|
|
source_indices (Any): Source indices value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
if source_indices is None:
|
|
return chunks
|
|
|
|
import mlx.core as mx
|
|
|
|
return mx.take(chunks, mx.array(source_indices, dtype=mx.int32), axis=1)
|
|
|
|
|
|
def _mlx_add_weighted_chunk(result, counter, chunk, window, start, length):
|
|
"""Implement the mlx add weighted chunk helper.
|
|
|
|
Args:
|
|
result (Any): Result value.
|
|
counter (Any): Counter value.
|
|
chunk (Any): Chunk value.
|
|
window (Any): Window value.
|
|
start (Any): Start value.
|
|
length (Any): Length value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
import mlx.core as mx
|
|
|
|
window = window[:length].astype(result.dtype)
|
|
weighted = chunk[..., :length].astype(result.dtype) * window
|
|
positions = mx.arange(start, start + length)
|
|
return result.at[:, :, positions].add(weighted), counter.at[:, :, positions].add(window)
|
|
|
|
|
|
def _mlx_finalize_overlap(result, counter, length_init, border):
|
|
"""Implement the mlx finalize overlap helper.
|
|
|
|
Args:
|
|
result (Any): Result value.
|
|
counter (Any): Counter value.
|
|
length_init (Any): Length init value.
|
|
border (Any): Border value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
import mlx.core as mx
|
|
|
|
estimated_sources = result / counter
|
|
if length_init > 2 * border and border > 0:
|
|
estimated_sources = estimated_sources[..., border:-border]
|
|
estimated_sources = np.array(estimated_sources, copy=False)
|
|
np.nan_to_num(estimated_sources, copy=False, nan=0.0)
|
|
return estimated_sources
|
|
|
|
|
|
def _can_demix_mlx_full(model, device):
|
|
"""Implement the can demix mlx full helper.
|
|
|
|
Args:
|
|
model (str): Model value.
|
|
device (Any): Device value.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
return (
|
|
torch.device(device).type == "mps"
|
|
and getattr(model, "mps_model_backend", None) == "mlx_full"
|
|
and hasattr(model, "mps_model_compute_dtype")
|
|
and hasattr(model, "mlx_forward_mx")
|
|
)
|
|
|
|
|
|
def demix_track_mlx_full(config, model, mix, device, pbar=False, source_indices=None, progress_callback=None):
|
|
"""Demix a tensor track with the full MLX inference path.
|
|
|
|
Args:
|
|
config (AttrDict | dict): Loaded pymss configuration.
|
|
model (str): Model value.
|
|
mix (np.ndarray): Mix value.
|
|
device (Any): Device value.
|
|
pbar (Any, optional): Pbar value. Defaults to False.
|
|
source_indices (Any, optional): Source indices value. Defaults to None.
|
|
progress_callback (Any, optional): Progress callback value. Defaults to None.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
import mlx.core as mx
|
|
|
|
C = config.audio.chunk_size
|
|
sample_rate = int(config.audio.get("sample_rate", 44100))
|
|
source_indices = _normalize_source_indices(config, source_indices)
|
|
step = _get_inference_step(config, C)
|
|
border = C - step
|
|
fade_size = min(C // 10, border)
|
|
batch_size = config.inference.batch_size
|
|
|
|
mix, length_init = _mlx_prepare_mix_for_chunks(mix, border)
|
|
starts, windows = _mlx_build_chunk_plan(mix.shape[1], C, step, fade_size)
|
|
result = mx.zeros((_source_count(config, source_indices), mix.shape[0], mix.shape[1]), dtype=mx.float32)
|
|
counter = mx.zeros((1, 1, mix.shape[1]), dtype=mx.float32)
|
|
progress = _ProgressContext(pbar, mix.shape[1], progress_callback, sample_rate=sample_rate)
|
|
|
|
for batch_start in range(0, len(starts), batch_size):
|
|
batch_indices = range(batch_start, min(batch_start + batch_size, len(starts)))
|
|
batch = [(_mlx_extract_chunk(mix, starts[idx], C), idx) for idx in batch_indices]
|
|
batch_count = len(batch)
|
|
chunks = _mlx_run_model_chunk(
|
|
model,
|
|
mx.stack([chunk for (chunk, _), _ in batch], axis=0),
|
|
C,
|
|
clear_cache_after_eval=bool(config.inference.get("mps_mlx_clear_cache", False)),
|
|
)
|
|
chunks = _mlx_select_sources(chunks, source_indices)
|
|
for j, ((_, length), idx) in enumerate(batch):
|
|
result, counter = _mlx_add_weighted_chunk(result, counter, chunks[j], windows[idx], starts[idx], length)
|
|
mx.eval(result, counter)
|
|
del chunks, batch
|
|
clear_mlx_cache()
|
|
progress.update(step * batch_count)
|
|
|
|
progress.close()
|
|
progress.emit(mix.shape[1])
|
|
estimated_sources = _mlx_finalize_overlap(result, counter, length_init, border)
|
|
sources = _sources_to_dict(config, estimated_sources, source_indices)
|
|
del result, counter, mix
|
|
clear_mlx_cache()
|
|
return sources
|
|
|
|
|
|
demix_track_mlx_roformer = demix_track_mlx_full
|
|
|
|
|
|
def demix_track(config, model, mix, device, pbar=False, source_indices=None, progress_callback=None):
|
|
"""Demix a tensor track with the PyTorch inference path.
|
|
|
|
Args:
|
|
config (AttrDict | dict): Loaded pymss configuration.
|
|
model (str): Model value.
|
|
mix (np.ndarray): Mix value.
|
|
device (Any): Device value.
|
|
pbar (Any, optional): Pbar value. Defaults to False.
|
|
source_indices (Any, optional): Source indices value. Defaults to None.
|
|
progress_callback (Any, optional): Progress callback value. Defaults to None.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
C = config.audio.chunk_size
|
|
sample_rate = int(config.audio.get("sample_rate", 44100))
|
|
source_indices = _normalize_source_indices(config, source_indices)
|
|
step = _get_inference_step(config, C)
|
|
border = C - step
|
|
fade_size = min(C // 10, border)
|
|
batch_size = config.inference.batch_size
|
|
|
|
mix, length_init = _prepare_mix_for_chunks(mix, border)
|
|
chunk_starts, chunk_windows = _build_chunk_plan(mix.shape[1], C, step, fade_size)
|
|
device_type = torch.device(device).type
|
|
use_complete_fast_path = device_type in ("cuda", "cpu")
|
|
mix_device = _model_mix(mix, device)
|
|
|
|
with _autocast(device, config.training.get("use_amp", True)):
|
|
with _inference_context(device):
|
|
result, counter = _init_overlap_buffers(config, mix, device, use_complete_fast_path, source_indices)
|
|
progress = _ProgressContext(pbar, mix.shape[1], progress_callback, sample_rate=sample_rate)
|
|
|
|
with _model_source_context(model, source_indices):
|
|
complete_chunks = 0
|
|
if use_complete_fast_path:
|
|
complete_chunks = _run_complete_chunks(
|
|
model,
|
|
mix_device,
|
|
chunk_windows,
|
|
result,
|
|
counter,
|
|
C,
|
|
step,
|
|
batch_size,
|
|
progress,
|
|
source_indices,
|
|
)
|
|
|
|
_run_tail_chunks(
|
|
model,
|
|
mix_device,
|
|
chunk_starts,
|
|
chunk_windows,
|
|
result,
|
|
counter,
|
|
C,
|
|
step,
|
|
batch_size,
|
|
complete_chunks,
|
|
progress,
|
|
source_indices,
|
|
)
|
|
progress.emit(mix.shape[1])
|
|
|
|
progress.close()
|
|
|
|
estimated_sources = _finalize_overlap(result, counter, length_init, border)
|
|
|
|
return _sources_to_dict(config, estimated_sources, source_indices)
|
|
|
|
|
|
def demix_track_demucs(config, model, mix, device, pbar=False, source_indices=None, progress_callback=None):
|
|
"""Demix a tensor track with Demucs-style inference.
|
|
|
|
Args:
|
|
config (AttrDict | dict): Loaded pymss configuration.
|
|
model (str): Model value.
|
|
mix (np.ndarray): Mix value.
|
|
device (Any): Device value.
|
|
pbar (Any, optional): Pbar value. Defaults to False.
|
|
source_indices (Any, optional): Source indices value. Defaults to None.
|
|
progress_callback (Any, optional): Progress callback value. Defaults to None.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
if _can_demix_mlx_full(model, device):
|
|
return demix_track_mlx_full(
|
|
config,
|
|
model,
|
|
mix.cpu().numpy(),
|
|
device,
|
|
pbar=pbar,
|
|
source_indices=source_indices,
|
|
progress_callback=progress_callback,
|
|
)
|
|
|
|
source_indices = _normalize_source_indices(config, source_indices)
|
|
source_names = _source_names(config)
|
|
S = len(source_names)
|
|
sample_rate = int(config.training.samplerate)
|
|
C = sample_rate * config.training.segment
|
|
batch_size = config.inference.batch_size
|
|
step = _get_inference_step(config, C)
|
|
|
|
with _autocast(device, config.training.get("use_amp", True)):
|
|
with _inference_context(device):
|
|
req_shape = (_source_count(config, source_indices),) + tuple(mix.shape)
|
|
result = torch.zeros(req_shape, dtype=torch.float32)
|
|
counter = torch.zeros(req_shape, dtype=torch.float32)
|
|
i = 0
|
|
batch_data = []
|
|
batch_locations = []
|
|
progress = _ProgressContext(pbar, mix.shape[1], progress_callback, sample_rate=sample_rate)
|
|
|
|
while i < mix.shape[1]:
|
|
part = mix[:, i : i + C].to(device)
|
|
length = part.shape[-1]
|
|
if length < C:
|
|
part = nn.functional.pad(input=part, pad=(0, C - length, 0, 0), mode="constant", value=0)
|
|
batch_data.append(part)
|
|
batch_locations.append((i, length))
|
|
i += step
|
|
|
|
if len(batch_data) >= batch_size or (i >= mix.shape[1]):
|
|
arr = torch.stack(batch_data, dim=0)
|
|
x = _select_sources(model(arr), source_indices)
|
|
for j, (start, l) in enumerate(batch_locations):
|
|
result[..., start : start + l] += x[j][..., :l].cpu()
|
|
counter[..., start : start + l] += 1.0
|
|
batch_data, batch_locations = [], []
|
|
|
|
progress.emit(min(i, mix.shape[1]))
|
|
|
|
progress.close()
|
|
progress.emit(mix.shape[1])
|
|
|
|
estimated_sources = (result / counter).cpu().numpy()
|
|
np.nan_to_num(estimated_sources, copy=False, nan=0.0)
|
|
|
|
if S == 1 and source_indices is None:
|
|
return estimated_sources
|
|
return _sources_to_dict(config, estimated_sources, source_indices)
|
|
|
|
|
|
def demix(
|
|
config, model, mix: NDArray, device, pbar=False, model_type: str = None, source_indices=None, progress_callback=None
|
|
) -> Dict[str, NDArray]:
|
|
"""Run chunked model inference and return separated sources.
|
|
|
|
Args:
|
|
config (AttrDict | dict): Loaded pymss configuration.
|
|
model (str): Model value.
|
|
mix (np.ndarray): Mix value.
|
|
device (Any): Device value.
|
|
pbar (Any, optional): Pbar value. Defaults to False.
|
|
model_type (Any, optional): Model type value. Defaults to None.
|
|
source_indices (Any, optional): Source indices value. Defaults to None.
|
|
progress_callback (Any, optional): Progress callback value. Defaults to None.
|
|
|
|
Returns:
|
|
Any: Computed result."""
|
|
if _can_demix_mlx_full(model, device):
|
|
return demix_track_mlx_full(
|
|
config, model, mix, device, pbar=pbar, source_indices=source_indices, progress_callback=progress_callback
|
|
)
|
|
mix = torch.tensor(mix, dtype=torch.float32)
|
|
if model_type in {"demucs", "tasnet", "legacy_demucs", "legacy_tasnet"}:
|
|
from .modules.legacy_demucs import apply_legacy_model
|
|
|
|
sample_rate = int(config.training.samplerate)
|
|
progress = _ProgressContext(
|
|
callback=progress_callback,
|
|
total=mix.shape[1],
|
|
sample_rate=sample_rate,
|
|
message="Processing audio",
|
|
)
|
|
progress.emit(0)
|
|
with _autocast(device, config.training.get("use_amp", True)):
|
|
with _inference_context(device):
|
|
estimates = (
|
|
apply_legacy_model(
|
|
model,
|
|
mix.to(device),
|
|
shifts=int(config.inference.get("shifts", 0)),
|
|
split=bool(config.inference.get("split", True)),
|
|
overlap=float(config.inference.get("overlap", 0.25)),
|
|
progress=pbar,
|
|
)
|
|
.cpu()
|
|
.numpy()
|
|
)
|
|
progress.emit(mix.shape[1])
|
|
return dict(zip(config.training.instruments, estimates))
|
|
if model_type == "htdemucs":
|
|
return demix_track_demucs(
|
|
config, model, mix, device, pbar=pbar, source_indices=source_indices, progress_callback=progress_callback
|
|
)
|
|
return demix_track(
|
|
config, model, mix, device, pbar=pbar, source_indices=source_indices, progress_callback=progress_callback
|
|
)
|