diff --git a/README.md b/README.md index ab0451e..e9fec25 100644 --- a/README.md +++ b/README.md @@ -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 ``` @@ -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` @@ -785,4 +785,4 @@ For questions or feedback, please raise an issue or reach out to @beveradb ([And - \ No newline at end of file + diff --git a/audio_separator/remote/api_client.py b/audio_separator/remote/api_client.py index 977b6e5..2c6536e 100644 --- a/audio_separator/remote/api_client.py +++ b/audio_separator/remote/api_client.py @@ -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). @@ -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 @@ -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: """ diff --git a/audio_separator/remote/cli.py b/audio_separator/remote/cli.py index 3ea625b..09d7818 100644 --- a/audio_separator/remote/cli.py +++ b/audio_separator/remote/cli.py @@ -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 diff --git a/audio_separator/remote/deploy_cloudrun.py b/audio_separator/remote/deploy_cloudrun.py index ccd2c8d..dab220a 100644 --- a/audio_separator/remote/deploy_cloudrun.py +++ b/audio_separator/remote/deploy_cloudrun.py @@ -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).""" @@ -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.""" diff --git a/audio_separator/remote/deploy_modal.py b/audio_separator/remote/deploy_modal.py index 3645c21..aa7f665 100755 --- a/audio_separator/remote/deploy_modal.py +++ b/audio_separator/remote/deploy_modal.py @@ -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: """ @@ -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: """ diff --git a/audio_separator/separator/architectures/mdxc_separator.py b/audio_separator/separator/architectures/mdxc_separator.py index 702d0a4..4c95d81 100644 --- a/audio_separator/separator/architectures/mdxc_separator.py +++ b/audio_separator/separator/architectures/mdxc_separator.py @@ -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. @@ -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) diff --git a/audio_separator/separator/separator.py b/audio_separator/separator/separator.py index 4561e0c..617252d 100644 --- a/audio_separator/separator/separator.py +++ b/audio_separator/separator/separator.py @@ -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 """ @@ -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}, ensemble_algorithm=None, ensemble_weights=None, ensemble_preset=None, diff --git a/audio_separator/utils/cli.py b/audio_separator/utils/cli.py index 23bbc90..6d26c8c 100755 --- a/audio_separator/utils/cli.py +++ b/audio_separator/utils/cli.py @@ -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() diff --git a/tests/unit/test_cli.py b/tests/unit/test_cli.py index ce49eef..a23cccc 100644 --- a/tests/unit/test_cli.py +++ b/tests/unit/test_cli.py @@ -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}, } diff --git a/tests/unit/test_mdxc_config.py b/tests/unit/test_mdxc_config.py new file mode 100644 index 0000000..1fd4b66 --- /dev/null +++ b/tests/unit/test_mdxc_config.py @@ -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) diff --git a/tests/unit/test_mdxc_roformer_chunking.py b/tests/unit/test_mdxc_roformer_chunking.py index 9d3ddd4..c33db80 100644 --- a/tests/unit/test_mdxc_roformer_chunking.py +++ b/tests/unit/test_mdxc_roformer_chunking.py @@ -9,7 +9,6 @@ from unittest.mock import Mock, MagicMock, patch import logging - class TestMDXCRoformerChunking: """Test cases for MDXC Roformer chunking and overlap functionality.""" @@ -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).""" diff --git a/tests/unit/test_remote_api_client.py b/tests/unit/test_remote_api_client.py index f9ef1a5..136a494 100644 --- a/tests/unit/test_remote_api_client.py +++ b/tests/unit/test_remote_api_client.py @@ -73,6 +73,22 @@ def test_separate_audio_success(self, mock_file, mock_post, api_client, mock_aud assert call_args[0][0] == "https://test-api.example.com/separate" assert "files" in call_args[1] assert "data" in call_args[1] + assert "mdxc_overlap" not in call_args[1]["data"] + assert "mdxc_batch_size" not in call_args[1]["data"] + + @patch("requests.Session.post") + @patch("builtins.open", new_callable=mock_open, read_data=b"fake audio content") + def test_separate_audio_with_mdxc_overrides(self, mock_file, mock_post, api_client, mock_audio_file): + mock_response = Mock() + mock_response.json.return_value = {"task_id": "test-task-123", "status": "submitted"} + mock_response.raise_for_status.return_value = None + mock_post.return_value = mock_response + + api_client.separate_audio(mock_audio_file, mdxc_overlap=4, mdxc_batch_size=2) + + data = mock_post.call_args[1]["data"] + assert data["mdxc_overlap"] == 4 + assert data["mdxc_batch_size"] == 2 @patch("requests.Session.post") @patch("builtins.open", new_callable=mock_open, read_data=b"fake audio content")