Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -520,8 +520,8 @@ Demucs Architecture Parameters:
MDXC Architecture Parameters:
--mdxc_segment_size MDXC_SEGMENT_SIZE Larger consumes more resources, but may give better results (default: 256). Example: --mdxc_segment_size=256
--mdxc_override_model_segment_size Override model default segment size instead of using the model default value. Example: --mdxc_override_model_segment_size
--mdxc_overlap MDXC_OVERLAP Amount of overlap between prediction windows, 2-50. Higher is better but slower (default: 8). Example: --mdxc_overlap=8
--mdxc_batch_size MDXC_BATCH_SIZE Larger consumes more RAM but may process slightly faster (default: 1). Example: --mdxc_batch_size=4
--mdxc_overlap MDXC_OVERLAP Number of overlapping prediction windows, 2-50. Higher is better but slower (default: model config, falling back to 8). Example: --mdxc_overlap=8
--mdxc_batch_size MDXC_BATCH_SIZE Larger consumes more RAM but may process slightly faster (default: model config, falling back to 1). Example: --mdxc_batch_size=4
--mdxc_pitch_shift MDXC_PITCH_SHIFT Shift audio pitch by a number of semitones while processing. May improve output for deep/high vocals. (default: 0). Example: --mdxc_pitch_shift=2
```

Expand Down Expand Up @@ -653,7 +653,7 @@ You can also rename specific stems:
- **`mdx_params`:** (Optional) MDX Architecture Specific Attributes & Defaults. `Default: {"hop_length": 1024, "segment_size": 256, "overlap": 0.25, "batch_size": 1, "enable_denoise": False}`
- **`vr_params`:** (Optional) VR Architecture Specific Attributes & Defaults. `Default: {"batch_size": 1, "window_size": 512, "aggression": 5, "enable_tta": False, "enable_post_process": False, "post_process_threshold": 0.2, "high_end_process": False}`
- **`demucs_params`:** (Optional) Demucs Architecture Specific Attributes & Defaults. `Default: {"segment_size": "Default", "shifts": 2, "overlap": 0.25, "segments_enabled": True}` _(Note: `segment_size` "Default" uses the model's internal default, typically 40 for older Demucs models and 10 for Demucs v4/htdemucs)_
- **`mdxc_params`:** (Optional) MDXC Architecture Specific Attributes & Defaults. `Default: {"segment_size": 256, "override_model_segment_size": False, "batch_size": 1, "overlap": 8, "pitch_shift": 0}`
- **`mdxc_params`:** (Optional) MDXC Architecture Specific Attributes & Defaults. `Default: {"segment_size": 256, "override_model_segment_size": False, "batch_size": None, "overlap": None, "pitch_shift": 0}` (`None` uses the model YAML's `inference` value, falling back to `1` for batch size and `8` for overlap.)
- **`ensemble_algorithm`:** (Optional) Algorithm to use for ensembling multiple models. `Default: 'avg_wave'`
- **`ensemble_weights`:** (Optional) Weights for each model in the ensemble. `Default: None` (equal weights)
- **`ensemble_preset`:** (Optional) Named ensemble preset (e.g. `'vocal_balanced'`, `'karaoke'`). Sets models, algorithm, and weights automatically. Use `Separator(info_only=True).list_ensemble_presets()` to see all. `Default: None`
Expand Down Expand Up @@ -785,4 +785,4 @@ For questions or feedback, please raise an issue or reach out to @beveradb ([And
<img src="https://contrib.rocks/image?repo=nomadkaraoke/python-audio-separator" />
</a>

</div>
</div>
15 changes: 9 additions & 6 deletions audio_separator/remote/api_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,8 +67,8 @@ def separate_audio(
# MDXC parameters
mdxc_segment_size: int = 256,
mdxc_override_model_segment_size: bool = False,
mdxc_overlap: int = 8,
mdxc_batch_size: int = 1,
mdxc_overlap: Optional[int] = None,
mdxc_batch_size: Optional[int] = None,
mdxc_pitch_shift: int = 0,
) -> dict:
"""Submit audio separation job (asynchronous processing).
Expand Down Expand Up @@ -133,12 +133,15 @@ def separate_audio(
# MDXC parameters
"mdxc_segment_size": mdxc_segment_size,
"mdxc_override_model_segment_size": mdxc_override_model_segment_size,
"mdxc_overlap": mdxc_overlap,
"mdxc_batch_size": mdxc_batch_size,
"mdxc_pitch_shift": mdxc_pitch_shift,
}
)

if mdxc_overlap is not None:
data["mdxc_overlap"] = mdxc_overlap
if mdxc_batch_size is not None:
data["mdxc_batch_size"] = mdxc_batch_size

# Add optional parameters only if they have non-default values
if output_bitrate:
data["output_bitrate"] = output_bitrate
Expand Down Expand Up @@ -209,8 +212,8 @@ def separate_audio_and_wait(
demucs_segments_enabled: bool = True,
mdxc_segment_size: int = 256,
mdxc_override_model_segment_size: bool = False,
mdxc_overlap: int = 8,
mdxc_batch_size: int = 1,
mdxc_overlap: Optional[int] = None,
mdxc_batch_size: Optional[int] = None,
mdxc_pitch_shift: int = 0,
) -> dict:
"""
Expand Down
4 changes: 2 additions & 2 deletions audio_separator/remote/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -137,8 +137,8 @@ def main():
mdxc_group = separate_parser.add_argument_group("MDXC Architecture Parameters")
mdxc_group.add_argument("--mdxc_segment_size", type=int, default=256, help="MDXC segment size (default: %(default)s)")
mdxc_group.add_argument("--mdxc_override_model_segment_size", action="store_true", help="Override MDXC model segment size")
mdxc_group.add_argument("--mdxc_overlap", type=int, default=8, help="MDXC overlap (default: %(default)s)")
mdxc_group.add_argument("--mdxc_batch_size", type=int, default=1, help="MDXC batch size (default: %(default)s)")
mdxc_group.add_argument("--mdxc_overlap", type=int, default=None, help="MDXC overlap (default: model config, falling back to 8)")
mdxc_group.add_argument("--mdxc_batch_size", type=int, default=None, help="MDXC batch size (default: model config, falling back to 1)")
mdxc_group.add_argument("--mdxc_pitch_shift", type=int, default=0, help="MDXC pitch shift (default: %(default)s)")

# Status command
Expand Down
8 changes: 4 additions & 4 deletions audio_separator/remote/deploy_cloudrun.py
Original file line number Diff line number Diff line change
Expand Up @@ -203,8 +203,8 @@ def separate_audio_sync(
# MDXC parameters
mdxc_segment_size: int = 256,
mdxc_override_model_segment_size: bool = False,
mdxc_overlap: int = 8,
mdxc_batch_size: int = 1,
mdxc_overlap: Optional[int] = None,
mdxc_batch_size: Optional[int] = None,
mdxc_pitch_shift: int = 0,
) -> dict:
"""Separate audio into stems. Runs synchronously (Cloud Run GPU handles one job at a time)."""
Expand Down Expand Up @@ -440,8 +440,8 @@ async def separate_audio(
# MDXC parameters
mdxc_segment_size: int = Form(256),
mdxc_override_model_segment_size: bool = Form(False),
mdxc_overlap: int = Form(8),
mdxc_batch_size: int = Form(1),
mdxc_overlap: Optional[int] = Form(None),
mdxc_batch_size: Optional[int] = Form(None),
mdxc_pitch_shift: int = Form(0),
) -> dict:
"""Upload an audio file (or provide a GCS URI) and separate it into stems."""
Expand Down
8 changes: 4 additions & 4 deletions audio_separator/remote/deploy_modal.py
Original file line number Diff line number Diff line change
Expand Up @@ -188,8 +188,8 @@ def separate_audio_function(
# MDXC parameters
mdxc_segment_size: int = 256,
mdxc_override_model_segment_size: bool = False,
mdxc_overlap: int = 8,
mdxc_batch_size: int = 1,
mdxc_overlap: Optional[int] = None,
mdxc_batch_size: Optional[int] = None,
mdxc_pitch_shift: int = 0,
) -> dict:
"""
Expand Down Expand Up @@ -575,8 +575,8 @@ async def separate_audio(
# MDXC parameters
mdxc_segment_size: int = Form(256, description="MDXC segment size"),
mdxc_override_model_segment_size: bool = Form(False, description="Override MDXC model segment size"),
mdxc_overlap: int = Form(8, description="MDXC overlap"),
mdxc_batch_size: int = Form(1, description="MDXC batch size"),
mdxc_overlap: Optional[int] = Form(None, description="MDXC overlap"),
mdxc_batch_size: Optional[int] = Form(None, description="MDXC batch size"),
mdxc_pitch_shift: int = Form(0, description="MDXC pitch shift"),
) -> dict:
"""
Expand Down
30 changes: 23 additions & 7 deletions audio_separator/separator/architectures/mdxc_separator.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,8 +93,26 @@ def __init__(self, common_config, arch_config):
# The segment size is set based on the value provided in a chosen model's associated config file (yaml).
self.override_model_segment_size = arch_config.get("override_model_segment_size", False)

self.overlap = arch_config.get("overlap", 8)
self.batch_size = arch_config.get("batch_size", 1)
inference_config = self.model_data.get("inference", {})
overlap = arch_config.get("overlap")
if overlap is None:
overlap = inference_config.get("num_overlap")
if overlap is None:
overlap = 8

batch_size = arch_config.get("batch_size")
if batch_size is None:
batch_size = inference_config.get("batch_size")
if batch_size is None:
batch_size = 1

if overlap <= 0:
raise ValueError("MDXC overlap must be greater than zero")
if batch_size <= 0:
raise ValueError("MDXC batch size must be greater than zero")

self.overlap = overlap
self.batch_size = batch_size

# Amount of pitch shift to apply during processing (this does NOT affect the pitch of the output audio):
# • Whole numbers indicate semitones.
Expand Down Expand Up @@ -354,11 +372,9 @@ def demix(self, mix: np.ndarray) -> dict:
f"Chunk size: {chunk_size} (using stft_hop_length={stft_hop_len} and dim_t={mdx_segment_size})"
)

# Align step to chunk_size by default for Roformer to avoid stride mismatches
# If a user-specified overlap (in seconds) results in a step larger than chunk_size, clamp it
desired_step = int(self.overlap * self.model_data_cfgdict.audio.sample_rate)
step = chunk_size if desired_step <= 0 else min(desired_step, chunk_size)
self.logger.debug(f"Step: {step} (desired={desired_step})")
# MDXC overlap is the number of overlapping prediction windows.
step = chunk_size // self.overlap
self.logger.debug(f"Step: {step} (overlap={self.overlap})")

# Create a weighting table and convert it to a PyTorch tensor
window = torch.tensor(signal.windows.hamming(chunk_size), dtype=torch.float32)
Expand Down
6 changes: 3 additions & 3 deletions audio_separator/separator/separator.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,8 +100,8 @@ class Separator:
MDXC Architecture Specific Attributes & Defaults:
segment_size: 256
override_model_segment_size: False
batch_size: 1
overlap: 8
batch_size: None (uses inference.batch_size from the model YAML, falling back to 1)
overlap: None (uses inference.num_overlap from the model YAML, falling back to 8)
pitch_shift: 0
"""

Expand All @@ -125,7 +125,7 @@ def __init__(
mdx_params={"hop_length": 1024, "segment_size": 256, "overlap": 0.25, "batch_size": 1, "enable_denoise": False},
vr_params={"batch_size": 1, "window_size": 512, "aggression": 5, "enable_tta": False, "enable_post_process": False, "post_process_threshold": 0.2, "high_end_process": False},
demucs_params={"segment_size": "Default", "shifts": 2, "overlap": 0.25, "segments_enabled": True},
mdxc_params={"segment_size": 256, "override_model_segment_size": False, "batch_size": 1, "overlap": 8, "pitch_shift": 0},
mdxc_params={"segment_size": 256, "override_model_segment_size": False, "batch_size": None, "overlap": None, "pitch_shift": 0},
Comment thread
coderabbitai[bot] marked this conversation as resolved.
ensemble_algorithm=None,
ensemble_weights=None,
ensemble_preset=None,
Expand Down
8 changes: 4 additions & 4 deletions audio_separator/utils/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -133,15 +133,15 @@ def main():

mdxc_segment_size_help = "Larger consumes more resources, but may give better results (default: %(default)s). Example: --mdxc_segment_size=256"
mdxc_override_model_segment_size_help = "Override model default segment size instead of using the model default value. Example: --mdxc_override_model_segment_size"
mdxc_overlap_help = "Amount of overlap between prediction windows, 2-50. Higher is better but slower (default: %(default)s). Example: --mdxc_overlap=8"
mdxc_batch_size_help = "Larger consumes more RAM but may process slightly faster (default: %(default)s). Example: --mdxc_batch_size=4"
mdxc_overlap_help = "Number of overlapping prediction windows, 2-50. Higher is better but slower (default: model config, falling back to 8). Example: --mdxc_overlap=8"
mdxc_batch_size_help = "Larger consumes more RAM but may process slightly faster (default: model config, falling back to 1). Example: --mdxc_batch_size=4"
mdxc_pitch_shift_help = "Shift audio pitch by a number of semitones while processing. May improve output for deep/high vocals. (default: %(default)s). Example: --mdxc_pitch_shift=2"

mdxc_params = parser.add_argument_group("MDXC Architecture Parameters")
mdxc_params.add_argument("--mdxc_segment_size", type=int, default=256, help=mdxc_segment_size_help)
mdxc_params.add_argument("--mdxc_override_model_segment_size", action="store_true", help=mdxc_override_model_segment_size_help)
mdxc_params.add_argument("--mdxc_overlap", type=int, default=8, help=mdxc_overlap_help)
mdxc_params.add_argument("--mdxc_batch_size", type=int, default=1, help=mdxc_batch_size_help)
mdxc_params.add_argument("--mdxc_overlap", type=int, default=None, help=mdxc_overlap_help)
mdxc_params.add_argument("--mdxc_batch_size", type=int, default=None, help=mdxc_batch_size_help)
mdxc_params.add_argument("--mdxc_pitch_shift", type=int, default=0, help=mdxc_pitch_shift_help)

args = parser.parse_args()
Expand Down
2 changes: 1 addition & 1 deletion tests/unit/test_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@ def common_expected_args():
"mdx_params": {"hop_length": 1024, "segment_size": 256, "overlap": 0.25, "batch_size": 1, "enable_denoise": False},
"vr_params": {"batch_size": 1, "window_size": 512, "aggression": 5, "enable_tta": False, "enable_post_process": False, "post_process_threshold": 0.2, "high_end_process": False},
"demucs_params": {"segment_size": "Default", "shifts": 2, "overlap": 0.25, "segments_enabled": True},
"mdxc_params": {"segment_size": 256, "batch_size": 1, "overlap": 8, "override_model_segment_size": False, "pitch_shift": 0},
"mdxc_params": {"segment_size": 256, "batch_size": None, "overlap": None, "override_model_segment_size": False, "pitch_shift": 0},
}


Expand Down
73 changes: 73 additions & 0 deletions tests/unit/test_mdxc_config.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
from unittest.mock import Mock, patch

import pytest
import torch
from ml_collections import ConfigDict

from audio_separator.separator.architectures.mdxc_separator import MDXCSeparator
from audio_separator.separator.common_separator import CommonSeparator


def _make_separator(arch_config, inference_config=None):
separator = MDXCSeparator.__new__(MDXCSeparator)
separator.logger = Mock()
separator.model_data = {
"inference": inference_config or {},
"training": {"target_instrument": "Vocals"},
}
separator.torch_device = torch.device("cpu")
separator.torch_device_cpu = torch.device("cpu")

def load_model():
separator.model_data_cfgdict = ConfigDict(separator.model_data)

with patch.object(CommonSeparator, "__init__", return_value=None), patch.object(MDXCSeparator, "load_model", side_effect=load_model):
MDXCSeparator.__init__(separator, {}, arch_config)

return separator


def test_mdxc_inference_defaults_come_from_model_config():
separator = _make_separator({"overlap": None, "batch_size": None}, {"num_overlap": 2, "batch_size": 4})

assert separator.overlap == 2
assert separator.batch_size == 4


def test_mdxc_explicit_inference_options_override_model_config():
separator = _make_separator({"overlap": 8, "batch_size": 1}, {"num_overlap": 2, "batch_size": 4})

assert separator.overlap == 8
assert separator.batch_size == 1


def test_mdxc_inference_defaults_fall_back_for_older_configs():
separator = _make_separator({"overlap": None, "batch_size": None})

assert separator.overlap == 8
assert separator.batch_size == 1


def test_mdxc_null_inference_values_use_fallbacks():
separator = _make_separator({"overlap": None, "batch_size": None}, {"num_overlap": None, "batch_size": None})

assert separator.overlap == 8
assert separator.batch_size == 1


@pytest.mark.parametrize(
("arch_config", "inference_config", "message"),
[
({"overlap": 0, "batch_size": 1}, {"num_overlap": 2}, "MDXC overlap"),
({"overlap": -1, "batch_size": 1}, {"num_overlap": 2}, "MDXC overlap"),
({"overlap": None, "batch_size": 1}, {"num_overlap": 0}, "MDXC overlap"),
({"overlap": None, "batch_size": 1}, {"num_overlap": -1}, "MDXC overlap"),
({"overlap": 2, "batch_size": 0}, {"batch_size": 1}, "MDXC batch size"),
({"overlap": 2, "batch_size": -1}, {"batch_size": 1}, "MDXC batch size"),
({"overlap": 2, "batch_size": None}, {"batch_size": 0}, "MDXC batch size"),
({"overlap": 2, "batch_size": None}, {"batch_size": -1}, "MDXC batch size"),
],
)
def test_mdxc_rejects_non_positive_inference_values(arch_config, inference_config, message):
with pytest.raises(ValueError, match=message):
_make_separator(arch_config, inference_config)
63 changes: 39 additions & 24 deletions tests/unit/test_mdxc_roformer_chunking.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,6 @@
from unittest.mock import Mock, MagicMock, patch
import logging


class TestMDXCRoformerChunking:
"""Test cases for MDXC Roformer chunking and overlap functionality."""

Expand Down Expand Up @@ -38,29 +37,45 @@ def test_chunk_size_falls_back_to_audio_hop_length(self):
# Test implementation for chunking optimization - placeholder for future implementation
pytest.skip("Chunking optimization not yet implemented")

def test_step_clamped_to_chunk_size(self):
"""T053: Step clamped to chunk_size (desired_step > chunk_size or ≤ 0)."""
chunk_size = 8192

# Test case 1: desired_step > chunk_size
desired_step_too_large = 10000
actual_step = min(desired_step_too_large, chunk_size)
assert actual_step == chunk_size

# Test case 2: desired_step ≤ 0
desired_step_zero = 0
actual_step = max(desired_step_zero, chunk_size // 4) # Use quarter chunk as minimum
assert actual_step == chunk_size // 4

# Test case 3: desired_step negative
desired_step_negative = -100
actual_step = max(desired_step_negative, chunk_size // 4)
assert actual_step == chunk_size // 4

# Test case 4: valid desired_step
desired_step_valid = 4096
actual_step = min(max(desired_step_valid, 1), chunk_size)
assert actual_step == desired_step_valid
def test_step_uses_overlap_divisor(self):
"""T053: Step follows the MSST num_overlap divisor semantics."""
from ml_collections import ConfigDict

from audio_separator.separator.architectures.mdxc_separator import MDXCSeparator

class IdentityModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.anchor = torch.nn.Parameter(torch.zeros(1))

def forward(self, audio):
return audio.unsqueeze(1)

separator = MDXCSeparator.__new__(MDXCSeparator)
separator.pitch_shift = 0
separator.is_roformer = True
separator.override_model_segment_size = False
separator.model_data_cfgdict = ConfigDict(
{
"inference": {"dim_t": 5},
"training": {"target_instrument": "Vocals", "instruments": ["Vocals"]},
"model": {"stft_hop_length": 2},
"audio": {"hop_length": 2},
}
)
separator.overlap = 2
separator.model_run = IdentityModel()
separator.is_primary_stem_main_target = False
separator.logger = Mock()
chunk_starts = []

with patch(
"audio_separator.separator.architectures.mdxc_separator.tqdm",
side_effect=lambda values: chunk_starts.extend(values) or values,
):
separator.demix(np.ones((2, 16), dtype=np.float32))

assert chunk_starts == [0, 4, 8, 12]

def test_overlap_add_short_output_safe(self):
"""T054: overlap_add handles shorter model output safely (safe_len)."""
Expand Down
Loading
Loading