import argparse import json import sys import warnings from .ensemble import ENSEMBLE_ALGORITHMS, save_ensemble_audio from .logger import get_separation_logger from .model_download import download_all, download_model from .model_registry import create_separator, list_models, resolve_model from .progress import _CliInferenceProgress from .workflow import load_workflow_file, run_workflow_file, validate_workflow, write_workflow_template warnings.filterwarnings("ignore", category=UserWarning) def _parse_key_value(values): """Parse key value. Args: values (Any): Values value. Returns: Any: Parsed value.""" result = {} for value in values or []: if "=" not in value: raise argparse.ArgumentTypeError(f"Expected key=value, got {value!r}") key, raw = value.split("=", 1) lowered = raw.lower() if lowered in {"true", "false"}: result[key] = lowered == "true" else: try: result[key] = int(raw) except ValueError: try: result[key] = float(raw) except ValueError: result[key] = raw return result def cmd_list(args): """Implement the cmd list helper. Args: args (argparse.Namespace): Parsed command-line arguments. Returns: Any: Computed result.""" rows = list_models(category=args.category, supported=None if args.all else True) if args.json: print(json.dumps([item.__dict__ for item in rows], ensure_ascii=False, indent=2)) return 0 for item in rows: status = "ok" if item.supported else item.unsupported_reason category = item.category_path or item.primary_category print(f"{item.name}\t{item.model_type or item.architecture}\t{category}\t{item.target_stem}\t{status}") return 0 def cmd_info(args): """Implement the cmd info helper. Args: args (argparse.Namespace): Parsed command-line arguments. Returns: Any: Computed result.""" resolved = resolve_model(args.model, model_dir=args.model_dir, require_supported=False, require_exists=False) entry = resolved["entry"] data = { "name": entry.name, "model_type": entry.model_type, "architecture": entry.architecture, "supported": entry.supported, "unsupported_reason": entry.unsupported_reason, "category": entry.category_path or entry.primary_category, "category_cn": " / ".join(filter(None, [entry.primary_category_cn, entry.secondary_category_cn])), "target_stem": entry.target_stem, "model_path": resolved["model_path"], "config_path": resolved["config_path"], "size_bytes": entry.size_bytes, } print(json.dumps(data, ensure_ascii=False, indent=2)) return 0 def cmd_download(args): """Implement the cmd download helper. Args: args (argparse.Namespace): Parsed command-line arguments. Returns: Any: Computed result.""" if args.model == "all": results = download_all( model_dir=args.model_dir, source=args.source, endpoint=args.endpoint, supported_only=args.supported_only, force=args.force, ) failed = [item for item in results if item.get("error")] print(f"Downloaded/skipped {len(results) - len(failed)} model(s), failed {len(failed)}.") for item in failed: print(f"ERROR {item['entry'].name}: {item['error']}", file=sys.stderr) return 1 if failed else 0 result = download_model( args.model, model_dir=args.model_dir, source=args.source, endpoint=args.endpoint, force=args.force, ) _print_download_result(result) return 0 def _print_download_result(result): """Print download result. Args: result (Any): Result value. Returns: None: This callable completes for its side effects.""" for path in result["skipped"]: print(f"exists {path}") for path in result["downloaded"]: print(f"downloaded {path}") def _ensure_model_files(args): """Ensure model files. Args: args (argparse.Namespace): Parsed command-line arguments. Returns: None: This callable completes for its side effects.""" try: resolve_model(args.model, model_dir=args.model_dir, require_supported=True, require_exists=True) except FileNotFoundError: result = download_model(args.model, model_dir=args.model_dir, source=args.source, endpoint=args.endpoint) _print_download_result(result) else: if args.download: result = download_model(args.model, model_dir=args.model_dir, source=args.source, endpoint=args.endpoint) _print_download_result(result) def cmd_infer(args): """Implement the cmd infer helper. Args: args (argparse.Namespace): Parsed command-line arguments. Returns: Any: Computed result.""" _ensure_model_files(args) logger = get_separation_logger() inference_progress = _CliInferenceProgress() with create_separator( args.model, model_dir=args.model_dir, device=args.device, device_ids=args.device_ids or [0], output_format=args.output_format, audio_params={ "wav_bit_depth": args.wav_bit_depth, "flac_bit_depth": args.flac_bit_depth, "mp3_bit_rate": args.mp3_bit_rate, "m4a_bit_rate": args.m4a_bit_rate, "m4a_aac_at_quality": args.m4a_aac_at_quality, }, use_tta=args.tta, store_dirs=args.output, save_as_folder=args.save_as_folder, logger=logger, debug=args.debug, progress_callback=inference_progress, inference_params=_parse_key_value(args.param), ) as separator: try: files = separator.process_folder(args.input) finally: inference_progress.close() logger.info(f"Processed {len(files)} file(s).") return 0 def cmd_ensemble(args): """Run the ensemble CLI command. Args: args (argparse.Namespace): Parsed command-line arguments. Returns: Any: Computed result.""" logger = get_separation_logger() output_path = save_ensemble_audio( args.files, args.output, algorithm=args.algorithm, weights=args.weights, output_format=args.output_format, audio_params={ "wav_bit_depth": args.wav_bit_depth, "flac_bit_depth": args.flac_bit_depth, "mp3_bit_rate": args.mp3_bit_rate, "m4a_bit_rate": args.m4a_bit_rate, "m4a_codec": args.m4a_codec, "m4a_aac_at_quality": args.m4a_aac_at_quality, }, logger=logger, ) logger.info(f"Saved ensemble audio to {output_path}") return 0 def cmd_workflow_init(args): """Write a starter workflow file.""" path = write_workflow_template(args.output, overwrite=args.force) print(f"Wrote workflow template to {path}") return 0 def cmd_workflow_validate(args): """Validate a workflow file without running inference.""" workflow = load_workflow_file(args.config) model_resolver = resolve_model if args.check_models or args.require_files else None validate_workflow( workflow, model_dir=args.model_dir, require_model_files=args.require_files, model_resolver=model_resolver, ) print(f"Workflow is valid: {len(workflow.steps)} step(s).") return 0 def cmd_workflow_run(args): """Run an audio workflow from a YAML/JSON file.""" logger = get_separation_logger() files = run_workflow_file( args.config, args.input, args.output, model_dir=args.model_dir, device=args.device, output_format=args.output_format, download=args.download, source=args.source, endpoint=args.endpoint, output_layout=args.output_layout, audio_params={ "wav_bit_depth": args.wav_bit_depth, "flac_bit_depth": args.flac_bit_depth, "mp3_bit_rate": args.mp3_bit_rate, "m4a_bit_rate": args.m4a_bit_rate, "m4a_codec": args.m4a_codec, "m4a_aac_at_quality": args.m4a_aac_at_quality, }, logger=logger, debug=args.debug, ) logger.info(f"Processed {len(files)} file(s).") return 0 def cmd_serve(args): """Implement the cmd serve helper. Args: args (argparse.Namespace): Parsed command-line arguments. Returns: Any: Computed result.""" from .server import ServerConfig, run_server config = ServerConfig( model=args.model, model_dir=args.model_dir, source=args.source, endpoint=args.endpoint, device=args.device, device_ids=args.device_ids or [0], api_key=args.api_key, host=args.host, port=args.port, debug=args.debug, inference_params=_parse_key_value(args.param), max_audio_seconds=args.max_audio_seconds, max_request_bytes=args.max_request_bytes, max_queue_size=args.max_queue_size, request_timeout_seconds=args.request_timeout_seconds, webui=args.webui, ) run_server(config) return 0 def build_parser(): """Build the pymss command-line parser. Args: None: This callable does not accept user-provided arguments. Returns: argparse.ArgumentParser: Configured CLI parser.""" parser = argparse.ArgumentParser( prog="pymss", description="Command-line interface for the pymss music source separation package.", formatter_class=lambda prog: argparse.RawTextHelpFormatter(prog, max_help_position=60), ) subparsers = parser.add_subparsers(dest="command", required=True) # ========================== # List models # ========================== list_parser = subparsers.add_parser( "list", help="List known models.", formatter_class=lambda prog: argparse.RawTextHelpFormatter(prog, max_help_position=60), ) list_parser.add_argument("--category", help="Filter by primary or secondary category.") list_parser.add_argument("--all", action="store_true", help="Include models that are not supported for inference yet.") list_parser.add_argument("--json", action="store_true") list_parser.set_defaults(func=cmd_list) # ========================== # Show model info # ========================== info_parser = subparsers.add_parser( "info", help="Show model metadata.", formatter_class=lambda prog: argparse.RawTextHelpFormatter(prog, max_help_position=60), ) info_parser.add_argument("model") info_parser.add_argument( "--model-dir", help="Local model cache directory. Defaults to PYMSS_MODEL_DIR, repository all_models if present, or ~/.cache/pymss/models.", ) info_parser.set_defaults(func=cmd_info) # ========================== # Download models # ========================== download_parser = subparsers.add_parser( "download", help="Download a model by name, or use 'all'.", formatter_class=lambda prog: argparse.RawTextHelpFormatter(prog, max_help_position=60), ) download_parser.add_argument("model") download_parser.add_argument( "--model-dir", help="Local model cache directory. Defaults to PYMSS_MODEL_DIR, repository all_models if present, or ~/.cache/pymss/models.", ) download_parser.add_argument("--source", default="modelscope", choices=["modelscope", "huggingface", "hf-mirror"]) download_parser.add_argument("--endpoint", help="Custom resolve endpoint. It must serve files by relative path.") download_parser.add_argument("--force", action="store_true") download_parser.add_argument("--supported-only", action="store_true", help="Only used with model='all'.") download_parser.set_defaults(func=cmd_download) # ========================== # Inference # ========================== infer_parser = subparsers.add_parser( "infer", help="Run inference by model name.", formatter_class=lambda prog: argparse.RawTextHelpFormatter(prog, max_help_position=60), ) infer_parser.add_argument("model") infer_parser.add_argument( "--model-dir", help="Local model cache directory. Defaults to PYMSS_MODEL_DIR, repository all_models if present, or ~/.cache/pymss/models.", ) infer_parser.add_argument( "--download", action="store_true", help="Check/download the model before inference. Missing model files are downloaded automatically.", ) infer_parser.add_argument("--source", default="modelscope", choices=["modelscope", "huggingface", "hf-mirror"]) infer_parser.add_argument("--endpoint", help="Custom resolve endpoint. It must serve files by relative path.") infer_parser.add_argument("-i", "--input", required=True, help="Input audio file or folder.") infer_parser.add_argument("-o", "--output", default="results", help="Output folder.") infer_parser.add_argument( "--save-as-folder", action="store_true", help="Save each input audio file's separated stems in a subfolder named after the audio file.", ) infer_parser.add_argument("--device", default="auto", choices=["auto", "cpu", "cuda", "mps", "mlx"]) infer_parser.add_argument( "--device-id", action="append", type=int, dest="device_ids", help="CUDA device id. Can be repeated." ) infer_parser.add_argument("--format", default="wav", choices=["wav", "flac", "mp3", "m4a"], dest="output_format") infer_parser.add_argument("--wav-bit-depth", default="FLOAT", choices=["FLOAT", "PCM_16", "PCM_24"]) infer_parser.add_argument("--flac-bit-depth", default="PCM_16", choices=["PCM_16", "PCM_24"]) infer_parser.add_argument("--mp3-bit-rate", default="320k") infer_parser.add_argument("--m4a-bit-rate", default="512k") infer_parser.add_argument("--m4a-codec", default="aac") infer_parser.add_argument("--m4a-aac-at-quality", default=2, type=int) infer_parser.add_argument("--tta", action="store_true", help="Enable test time augmentation.") infer_parser.add_argument("--debug", action="store_true") infer_parser.add_argument( "--param", action="append", default=[], help="Inference override as key=value, for example --param batch_size=2." ) infer_parser.set_defaults(func=cmd_infer) # ========================== # Ensemble # ========================== ensemble_parser = subparsers.add_parser( "ensemble", help="Combine multiple audio files with an ensemble algorithm.", formatter_class=lambda prog: argparse.RawTextHelpFormatter(prog, max_help_position=60), ) ensemble_parser.add_argument("files", nargs="+", help="Input audio files. At least two files are required.") ensemble_parser.add_argument( "-a", "--algorithm", default="avg_wave", choices=ENSEMBLE_ALGORITHMS, help="Ensemble algorithm.", ) ensemble_parser.add_argument( "-w", "--weights", nargs="+", type=float, help="Input weights, for example --weights 1 0.8 1.2. Defaults to all 1.", ) ensemble_parser.add_argument("-o", "--output", required=True, help="Output audio file.") ensemble_parser.add_argument("--format", choices=["wav", "flac", "mp3", "m4a"], dest="output_format") ensemble_parser.add_argument("--wav-bit-depth", default="FLOAT", choices=["FLOAT", "PCM_16", "PCM_24"]) ensemble_parser.add_argument("--flac-bit-depth", default="PCM_16", choices=["PCM_16", "PCM_24"]) ensemble_parser.add_argument("--mp3-bit-rate", default="320k") ensemble_parser.add_argument("--m4a-bit-rate", default="512k") ensemble_parser.add_argument("--m4a-codec", default="aac") ensemble_parser.add_argument("--m4a-aac-at-quality", default=2, type=int) ensemble_parser.set_defaults(func=cmd_ensemble) # ========================== # Workflow # ========================== workflow_parser = subparsers.add_parser( "workflow", help="Create, validate, or run an automatic multi-model workflow.", formatter_class=lambda prog: argparse.RawTextHelpFormatter(prog, max_help_position=60), ) workflow_subparsers = workflow_parser.add_subparsers(dest="workflow_command", required=True) workflow_init_parser = workflow_subparsers.add_parser( "init", help="Write a starter workflow YAML file.", formatter_class=lambda prog: argparse.RawTextHelpFormatter(prog, max_help_position=60), ) workflow_init_parser.add_argument("-o", "--output", default="workflow.yaml", help="Workflow file to create.") workflow_init_parser.add_argument("--force", action="store_true", help="Overwrite the output file if it exists.") workflow_init_parser.set_defaults(func=cmd_workflow_init) workflow_validate_parser = workflow_subparsers.add_parser( "validate", help="Validate a workflow YAML/JSON file.", formatter_class=lambda prog: argparse.RawTextHelpFormatter(prog, max_help_position=60), ) workflow_validate_parser.add_argument("-c", "--config", required=True, help="Workflow YAML/JSON file.") workflow_validate_parser.add_argument( "--model-dir", help="Local model cache directory used when --require-files is set.", ) workflow_validate_parser.add_argument( "--check-models", action="store_true", help="Also check that every referenced model exists in the catalog.", ) workflow_validate_parser.add_argument( "--require-files", action="store_true", help="Also require every referenced catalog model file to exist locally.", ) workflow_validate_parser.set_defaults(func=cmd_workflow_validate) workflow_run_parser = workflow_subparsers.add_parser( "run", help="Run inference through a workflow YAML/JSON file.", formatter_class=lambda prog: argparse.RawTextHelpFormatter(prog, max_help_position=60), ) workflow_run_parser.add_argument("-c", "--config", required=True, help="Workflow YAML/JSON file.") workflow_run_parser.add_argument("-i", "--input", required=True, help="Input audio file or folder.") workflow_run_parser.add_argument("-o", "--output", default="results", help="Output folder.") workflow_run_parser.add_argument( "--output-layout", default="folders", choices=["folders", "flat"], help=( "Workflow output layout. 'folders' keeps each input under /