mirror of
https://github.com/RVC-Project/Retrieval-based-Voice-Conversion-WebUI.git
synced 2026-09-01 11:38:32 +02:00
716 lines
25 KiB
Python
716 lines
25 KiB
Python
from __future__ import annotations
|
|
|
|
import os
|
|
import re
|
|
from dataclasses import dataclass, field
|
|
from pathlib import Path
|
|
from typing import Any, Callable
|
|
|
|
import numpy as np
|
|
import yaml
|
|
|
|
|
|
WORKFLOW_TEMPLATE = """version: 1
|
|
|
|
defaults:
|
|
device: auto
|
|
output_format: wav
|
|
model_dir: null
|
|
inference_params:
|
|
normalize: false
|
|
|
|
steps:
|
|
- id: split
|
|
model: bs_roformer_voc_hyperacev2
|
|
input: input
|
|
stems: [vocals, other]
|
|
inference_params:
|
|
overlap_size: 48000
|
|
save:
|
|
vocals: vocal
|
|
other: other
|
|
|
|
- id: dereverb
|
|
model: UVR-DeReverb-aufr33-jarredou_4band_v4_ms_fullband
|
|
input: split.other
|
|
stems: [Dry]
|
|
inference_params:
|
|
overlap_size: 22050
|
|
save:
|
|
Dry: dry
|
|
|
|
- id: harmony
|
|
model: your_harmony_model
|
|
input: dereverb.Dry
|
|
stems: [other]
|
|
inference_params:
|
|
overlap_size: 22050
|
|
save:
|
|
other: harmony_other
|
|
"""
|
|
|
|
|
|
_STEP_ID_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_-]*$")
|
|
_OUTPUT_FORMATS = {"wav", "flac", "mp3", "m4a"}
|
|
_OUTPUT_LAYOUTS = {"folders", "flat"}
|
|
_DEFAULT_AUDIO_PARAMS = {
|
|
"wav_bit_depth": "FLOAT",
|
|
"flac_bit_depth": "PCM_24",
|
|
"mp3_bit_rate": "320k",
|
|
"m4a_bit_rate": "512k",
|
|
"m4a_codec": "aac",
|
|
"m4a_aac_at_quality": 2,
|
|
}
|
|
|
|
|
|
class WorkflowError(ValueError):
|
|
"""Raised when a workflow definition or run is invalid."""
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class WorkflowStep:
|
|
id: str
|
|
model: str | None = None
|
|
input: str = "input"
|
|
stems: list[str] | None = None
|
|
save: dict[str, Any] = field(default_factory=dict)
|
|
model_type: str | None = None
|
|
model_path: str | None = None
|
|
config_path: str | None = None
|
|
device: str | None = None
|
|
model_dir: str | None = None
|
|
output_format: str | None = None
|
|
inference_params: dict[str, Any] = field(default_factory=dict)
|
|
use_tta: bool | None = None
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class Workflow:
|
|
version: int
|
|
defaults: dict[str, Any]
|
|
steps: list[WorkflowStep]
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class AudioArtifact:
|
|
audio: np.ndarray
|
|
sample_rate: int
|
|
|
|
|
|
@dataclass
|
|
class WorkflowTrackState:
|
|
path: str
|
|
track_name: str
|
|
artifacts: dict[str, AudioArtifact] = field(default_factory=dict)
|
|
active: bool = True
|
|
|
|
|
|
def load_workflow_file(path: str | os.PathLike) -> Workflow:
|
|
"""Load a workflow YAML/JSON file."""
|
|
workflow_path = Path(path)
|
|
try:
|
|
data = yaml.safe_load(workflow_path.read_text(encoding="utf-8"))
|
|
except yaml.YAMLError as exc:
|
|
raise WorkflowError(f"Invalid workflow YAML: {exc}") from exc
|
|
except OSError as exc:
|
|
raise WorkflowError(f"Cannot read workflow file: {workflow_path}") from exc
|
|
return load_workflow_data(data)
|
|
|
|
|
|
def load_workflow_data(data: Any) -> Workflow:
|
|
"""Parse workflow data from a Python mapping."""
|
|
if not isinstance(data, dict):
|
|
raise WorkflowError("Workflow file must contain a mapping.")
|
|
version = data.get("version")
|
|
if version != 1:
|
|
raise WorkflowError("workflow version must be 1.")
|
|
defaults = data.get("defaults") or {}
|
|
if not isinstance(defaults, dict):
|
|
raise WorkflowError("defaults must be a mapping.")
|
|
raw_steps = data.get("steps")
|
|
if not isinstance(raw_steps, list) or not raw_steps:
|
|
raise WorkflowError("steps must be a non-empty list.")
|
|
steps = [_parse_step(index, item) for index, item in enumerate(raw_steps, start=1)]
|
|
workflow = Workflow(version=int(version), defaults=dict(defaults), steps=steps)
|
|
validate_workflow_structure(workflow)
|
|
return workflow
|
|
|
|
|
|
def write_workflow_template(path: str | os.PathLike, *, overwrite: bool = False) -> Path:
|
|
"""Write a starter workflow YAML file."""
|
|
output_path = Path(path)
|
|
if output_path.exists() and not overwrite:
|
|
raise WorkflowError(f"Workflow file already exists: {output_path}")
|
|
output_path.parent.mkdir(parents=True, exist_ok=True)
|
|
output_path.write_text(WORKFLOW_TEMPLATE, encoding="utf-8")
|
|
return output_path
|
|
|
|
|
|
def validate_workflow(
|
|
workflow: Workflow,
|
|
*,
|
|
model_dir: str | os.PathLike | None = None,
|
|
require_model_files: bool = False,
|
|
model_resolver: Callable[..., Any] | None = None,
|
|
) -> Workflow:
|
|
"""Validate workflow references and optionally model catalog entries."""
|
|
validate_workflow_structure(workflow)
|
|
_validate_step_references(workflow)
|
|
if model_resolver is not None or require_model_files:
|
|
resolver = model_resolver or _default_model_resolver
|
|
for step in workflow.steps:
|
|
if step.model_path:
|
|
_validate_explicit_model_files(step, require_model_files=require_model_files)
|
|
continue
|
|
step_model_dir = _step_option(workflow, step, "model_dir", model_dir)
|
|
resolver(
|
|
step.model,
|
|
model_dir=step_model_dir,
|
|
require_supported=True,
|
|
require_exists=require_model_files,
|
|
)
|
|
return workflow
|
|
|
|
|
|
def validate_workflow_structure(workflow: Workflow) -> Workflow:
|
|
"""Validate syntax that does not require cross-step analysis."""
|
|
seen = set()
|
|
for step in workflow.steps:
|
|
if not _STEP_ID_RE.match(step.id):
|
|
raise WorkflowError(f"Invalid step id {step.id!r}.")
|
|
if step.id in seen:
|
|
raise WorkflowError(f"Duplicate step id: {step.id}")
|
|
seen.add(step.id)
|
|
if bool(step.model) == bool(step.model_path):
|
|
raise WorkflowError(f"step {step.id!r} requires exactly one of model or model_path.")
|
|
if step.model_path and not step.model_type:
|
|
raise WorkflowError(f"step {step.id!r} requires model_type when model_path is used.")
|
|
if step.stems is not None and not step.stems:
|
|
raise WorkflowError(f"step {step.id!r} stems must not be empty.")
|
|
if step.output_format is not None and str(step.output_format).lower() not in _OUTPUT_FORMATS:
|
|
raise WorkflowError(f"step {step.id!r} has unsupported output_format {step.output_format!r}.")
|
|
default_format = workflow.defaults.get("output_format")
|
|
if default_format is not None and str(default_format).lower() not in _OUTPUT_FORMATS:
|
|
raise WorkflowError(f"defaults.output_format must be one of: {sorted(_OUTPUT_FORMATS)}.")
|
|
default_inference_params = workflow.defaults.get("inference_params")
|
|
if default_inference_params is not None and not isinstance(default_inference_params, dict):
|
|
raise WorkflowError("defaults.inference_params must be a mapping.")
|
|
return workflow
|
|
|
|
|
|
class WorkflowRunner:
|
|
"""Run a parsed pymss workflow over one file or a direct folder."""
|
|
|
|
def __init__(
|
|
self,
|
|
workflow: Workflow,
|
|
*,
|
|
model_dir: str | os.PathLike | None = None,
|
|
device: str | None = None,
|
|
output_format: str | None = None,
|
|
download: bool = False,
|
|
source: str = "modelscope",
|
|
endpoint: str | None = None,
|
|
audio_params: dict[str, Any] | None = None,
|
|
logger: Any = None,
|
|
debug: bool = False,
|
|
separator_factory: Callable[..., Any] | None = None,
|
|
audio_loader: Callable[..., Any] | None = None,
|
|
audio_saver: Callable[..., Any] | None = None,
|
|
continue_on_error: bool = False,
|
|
output_layout: str = "folders",
|
|
):
|
|
self.workflow = validate_workflow(workflow)
|
|
self.model_dir = model_dir
|
|
self.device = device
|
|
self.output_format = output_format
|
|
self.download = bool(download)
|
|
self.source = source
|
|
self.endpoint = endpoint
|
|
self.audio_params = {**_DEFAULT_AUDIO_PARAMS, **(audio_params or {})}
|
|
self.logger = logger
|
|
self.debug = bool(debug)
|
|
self.separator_factory = separator_factory or _default_separator_factory
|
|
self.audio_loader = audio_loader or _default_audio_loader
|
|
self.audio_saver = audio_saver or _default_audio_saver
|
|
self.continue_on_error = bool(continue_on_error)
|
|
self.output_layout = _validate_output_layout(output_layout)
|
|
|
|
def run(self, input_path: str | os.PathLike, output_dir: str | os.PathLike) -> list[str]:
|
|
"""Run the workflow and return successfully processed basenames."""
|
|
paths = _input_files(input_path)
|
|
output_root = Path(output_dir)
|
|
tracks = [
|
|
track
|
|
for path, track_name in zip(paths, _unique_track_names(paths))
|
|
for track in [self._load_track(path, track_name)]
|
|
if track is not None
|
|
]
|
|
for step in self.workflow.steps:
|
|
active_tracks = [track for track in tracks if track.active]
|
|
if not active_tracks:
|
|
break
|
|
try:
|
|
with self._open_separator(step) as separator:
|
|
for track in active_tracks:
|
|
self._run_step_for_track(step, separator, track, output_root)
|
|
except Exception as exc:
|
|
if not self.continue_on_error:
|
|
raise
|
|
for track in active_tracks:
|
|
self._mark_track_failed(track, exc)
|
|
return [os.path.basename(track.path) for track in tracks if track.active]
|
|
|
|
def _load_track(self, path: str, track_name: str) -> WorkflowTrackState | None:
|
|
try:
|
|
mix, sr = self.audio_loader(path, sr=None, mono=False)
|
|
return WorkflowTrackState(
|
|
path=path,
|
|
track_name=track_name,
|
|
artifacts={"input": AudioArtifact(_to_model_audio(mix), int(sr))},
|
|
)
|
|
except Exception as exc:
|
|
if self.continue_on_error and self.logger is not None:
|
|
self.logger.warning("Cannot process workflow track %s: %s", path, exc)
|
|
return None
|
|
raise
|
|
|
|
def _run_step_for_track(
|
|
self,
|
|
step: WorkflowStep,
|
|
separator: Any,
|
|
track: WorkflowTrackState,
|
|
output_root: Path,
|
|
) -> None:
|
|
try:
|
|
artifact = _resolve_input_artifact(track.artifacts, step.input)
|
|
sample_rate = int(separator.config.audio.get("sample_rate", artifact.sample_rate))
|
|
model_audio = _ensure_sample_rate(_to_model_audio(artifact.audio), artifact.sample_rate, sample_rate)
|
|
stems = _requested_stems(step)
|
|
if getattr(separator, "model_type", None) == "vr":
|
|
results = separator.separate(model_audio, pbar=False)
|
|
else:
|
|
results = separator.separate(model_audio, pbar=False, stems=stems)
|
|
selected = _select_results(step, results)
|
|
for stem, audio in selected.items():
|
|
track.artifacts[f"{step.id}.{stem}"] = AudioArtifact(_to_model_audio(audio), sample_rate)
|
|
self._save_results(step, selected, sample_rate, output_root, track.track_name)
|
|
del selected, results
|
|
except Exception as exc:
|
|
if not self.continue_on_error:
|
|
raise
|
|
self._mark_track_failed(track, exc)
|
|
|
|
def _mark_track_failed(self, track: WorkflowTrackState, exc: Exception) -> None:
|
|
track.active = False
|
|
if self.logger is not None:
|
|
self.logger.warning("Cannot process workflow track %s: %s", track.path, exc)
|
|
|
|
def _open_separator(self, step: WorkflowStep):
|
|
if self.download and step.model:
|
|
from .model_download import download_model
|
|
|
|
download_model(
|
|
step.model,
|
|
model_dir=_step_option(self.workflow, step, "model_dir", self.model_dir),
|
|
source=self.source,
|
|
endpoint=self.endpoint,
|
|
)
|
|
separator_kwargs = {
|
|
"model_dir": _step_option(self.workflow, step, "model_dir", self.model_dir),
|
|
"device": _step_option(self.workflow, step, "device", self.device),
|
|
"output_format": _step_option(self.workflow, step, "output_format", self.output_format) or "wav",
|
|
"audio_params": self.audio_params,
|
|
"use_tta": bool(_step_option(self.workflow, step, "use_tta", None) or False),
|
|
"logger": self.logger,
|
|
"debug": self.debug,
|
|
"inference_params": _merged_inference_params(self.workflow, step),
|
|
}
|
|
model_name = step.model
|
|
if step.model_path:
|
|
model_name = Path(step.model_path).stem
|
|
separator_kwargs.update(
|
|
{
|
|
"model_type": step.model_type,
|
|
"model_path": step.model_path,
|
|
"config_path": step.config_path,
|
|
}
|
|
)
|
|
separator = self.separator_factory(
|
|
model_name,
|
|
**separator_kwargs,
|
|
)
|
|
return _SeparatorContext(separator)
|
|
|
|
def _save_results(
|
|
self,
|
|
step: WorkflowStep,
|
|
results: dict[str, np.ndarray],
|
|
sample_rate: int,
|
|
output_root: Path,
|
|
track_name: str,
|
|
) -> None:
|
|
output_format = str(_step_option(self.workflow, step, "output_format", self.output_format) or "wav").lower()
|
|
for stem, audio in results.items():
|
|
save_dirs = _save_dirs(step, stem)
|
|
for save_dir in save_dirs:
|
|
target_dir = output_root / save_dir
|
|
if self.output_layout == "folders":
|
|
target_dir = output_root / track_name / save_dir
|
|
target_dir.mkdir(parents=True, exist_ok=True)
|
|
safe_stem = _safe_filename_part(stem)
|
|
target = target_dir / f"{track_name}_{safe_stem}.{output_format}"
|
|
self.audio_saver(str(target), _to_save_audio(audio), sample_rate, output_format, self.audio_params)
|
|
|
|
|
|
class _SeparatorContext:
|
|
def __init__(self, separator):
|
|
self.separator = separator
|
|
|
|
def __enter__(self):
|
|
enter = getattr(self.separator, "__enter__", None)
|
|
return enter() if enter is not None else self.separator
|
|
|
|
def __exit__(self, exc_type, exc_value, traceback):
|
|
exit_method = getattr(self.separator, "__exit__", None)
|
|
if exit_method is not None:
|
|
return exit_method(exc_type, exc_value, traceback)
|
|
close = getattr(self.separator, "close", None)
|
|
if close is not None:
|
|
close()
|
|
return False
|
|
|
|
|
|
def run_workflow_file(
|
|
config_path: str | os.PathLike,
|
|
input_path: str | os.PathLike,
|
|
output_dir: str | os.PathLike,
|
|
**runner_kwargs,
|
|
) -> list[str]:
|
|
"""Load and run a workflow file."""
|
|
workflow = load_workflow_file(config_path)
|
|
return WorkflowRunner(workflow, **runner_kwargs).run(input_path, output_dir)
|
|
|
|
|
|
def _validate_output_layout(value: str) -> str:
|
|
layout = str(value).strip().lower()
|
|
if layout not in _OUTPUT_LAYOUTS:
|
|
raise WorkflowError(f"output_layout must be one of: {sorted(_OUTPUT_LAYOUTS)}.")
|
|
return layout
|
|
|
|
|
|
def _parse_step(index: int, data: Any) -> WorkflowStep:
|
|
if not isinstance(data, dict):
|
|
raise WorkflowError(f"step #{index} must be a mapping.")
|
|
step_id = data.get("id")
|
|
if not isinstance(step_id, str) or not step_id.strip():
|
|
raise WorkflowError(f"step #{index} requires a non-empty id.")
|
|
model = _parse_optional_string(data.get("model"))
|
|
model_path = _parse_optional_string(data.get("model_path"))
|
|
return WorkflowStep(
|
|
id=step_id.strip(),
|
|
model=model,
|
|
input=_parse_input_value(data.get("input", "input"), step_id),
|
|
stems=_parse_stems(data.get("stems"), step_id),
|
|
save=_parse_save(data.get("save"), step_id),
|
|
model_type=_parse_optional_string(data.get("model_type")),
|
|
model_path=model_path,
|
|
config_path=_parse_optional_string(data.get("config_path")),
|
|
device=_parse_optional_string(data.get("device")),
|
|
model_dir=_parse_optional_string(data.get("model_dir")),
|
|
output_format=_parse_optional_string(data.get("output_format")),
|
|
inference_params=_parse_mapping(data.get("inference_params"), step_id, "inference_params"),
|
|
use_tta=_parse_optional_bool(data.get("use_tta"), step_id, "use_tta"),
|
|
)
|
|
|
|
|
|
def _parse_input_value(value: Any, step_id: str) -> str:
|
|
if not isinstance(value, str) or not value.strip():
|
|
raise WorkflowError(f"step {step_id!r} input must be a non-empty string.")
|
|
return value.strip()
|
|
|
|
|
|
def _parse_stems(value: Any, step_id: str) -> list[str] | None:
|
|
if value is None:
|
|
return None
|
|
if isinstance(value, str):
|
|
stems = [value]
|
|
elif isinstance(value, list):
|
|
stems = value
|
|
else:
|
|
raise WorkflowError(f"step {step_id!r} stems must be a string or list.")
|
|
result = [str(item).strip() for item in stems if str(item).strip()]
|
|
if not result:
|
|
raise WorkflowError(f"step {step_id!r} stems must not be empty.")
|
|
return result
|
|
|
|
|
|
def _parse_save(value: Any, step_id: str) -> dict[str, Any]:
|
|
if value is None:
|
|
return {}
|
|
if not isinstance(value, dict):
|
|
raise WorkflowError(f"step {step_id!r} save must be a mapping.")
|
|
result = {}
|
|
for stem, target in value.items():
|
|
stem_name = str(stem).strip()
|
|
if not stem_name:
|
|
raise WorkflowError(f"step {step_id!r} save contains an empty stem name.")
|
|
result[stem_name] = target
|
|
return result
|
|
|
|
|
|
def _parse_mapping(value: Any, step_id: str, field_name: str) -> dict[str, Any]:
|
|
if value is None:
|
|
return {}
|
|
if not isinstance(value, dict):
|
|
raise WorkflowError(f"step {step_id!r} {field_name} must be a mapping.")
|
|
return dict(value)
|
|
|
|
|
|
def _parse_optional_string(value: Any) -> str | None:
|
|
if value is None:
|
|
return None
|
|
value = str(value).strip()
|
|
return value or None
|
|
|
|
|
|
def _parse_optional_bool(value: Any, step_id: str, field_name: str) -> bool | None:
|
|
if value is None:
|
|
return None
|
|
if isinstance(value, bool):
|
|
return value
|
|
raise WorkflowError(f"step {step_id!r} {field_name} must be a boolean.")
|
|
|
|
|
|
def _validate_step_references(workflow: Workflow) -> None:
|
|
seen = {"input"}
|
|
required_outputs: dict[str, set[str]] = {step.id: set() for step in workflow.steps}
|
|
for step in workflow.steps:
|
|
if step.input != "input":
|
|
ref_step, ref_stem = _split_artifact_ref(step.input, step.id)
|
|
if ref_step not in seen:
|
|
raise WorkflowError(f"step {step.id!r} input references unknown step: {ref_step}")
|
|
required_outputs.setdefault(ref_step, set()).add(ref_stem)
|
|
for stem in step.save:
|
|
required_outputs[step.id].add(stem)
|
|
seen.add(step.id)
|
|
|
|
for step in workflow.steps:
|
|
if step.stems is None:
|
|
continue
|
|
requested = {stem.lower() for stem in step.stems}
|
|
for stem in required_outputs.get(step.id, set()):
|
|
if stem.lower() not in requested:
|
|
raise WorkflowError(
|
|
f"step {step.id!r} must request {step.id}.{stem}; add {stem!r} to stems or omit stems."
|
|
)
|
|
|
|
|
|
def _split_artifact_ref(value: str, current_step_id: str) -> tuple[str, str]:
|
|
if "." not in value:
|
|
raise WorkflowError(f"step {current_step_id!r} input must be 'input' or '<step>.<stem>'.")
|
|
step_id, stem = value.split(".", 1)
|
|
step_id = step_id.strip()
|
|
stem = stem.strip()
|
|
if not step_id or not stem:
|
|
raise WorkflowError(f"step {current_step_id!r} input must be 'input' or '<step>.<stem>'.")
|
|
return step_id, stem
|
|
|
|
|
|
def _validate_explicit_model_files(step: WorkflowStep, *, require_model_files: bool) -> None:
|
|
if not require_model_files:
|
|
return
|
|
missing = []
|
|
if step.model_path and not Path(step.model_path).is_file():
|
|
missing.append(step.model_path)
|
|
if step.config_path and not Path(step.config_path).is_file():
|
|
missing.append(step.config_path)
|
|
if missing:
|
|
raise FileNotFoundError("Missing model file(s): " + ", ".join(missing))
|
|
|
|
|
|
def _resolve_input_artifact(artifacts: dict[str, AudioArtifact], ref: str) -> AudioArtifact:
|
|
if ref == "input":
|
|
return artifacts["input"]
|
|
if ref in artifacts:
|
|
return artifacts[ref]
|
|
ref_step, ref_stem = ref.split(".", 1)
|
|
for key, artifact in artifacts.items():
|
|
if not key.startswith(f"{ref_step}."):
|
|
continue
|
|
_, stem = key.split(".", 1)
|
|
if stem.lower() == ref_stem.lower():
|
|
return artifact
|
|
raise WorkflowError(f"Missing workflow input artifact: {ref}")
|
|
|
|
|
|
def _requested_stems(step: WorkflowStep) -> list[str] | None:
|
|
if step.stems is not None:
|
|
return list(step.stems)
|
|
if step.save:
|
|
return list(step.save)
|
|
return None
|
|
|
|
|
|
def _select_results(step: WorkflowStep, results: dict[str, Any]) -> dict[str, np.ndarray]:
|
|
requested = _requested_stems(step)
|
|
if requested is None:
|
|
requested = list(results)
|
|
selected = {}
|
|
for stem in requested:
|
|
actual = _find_stem(results, stem)
|
|
selected[actual] = np.asarray(results[actual], dtype=np.float32)
|
|
return selected
|
|
|
|
|
|
def _find_stem(results: dict[str, Any], stem: str) -> str:
|
|
if stem in results:
|
|
return stem
|
|
lower = str(stem).lower()
|
|
for key in results:
|
|
if str(key).lower() == lower:
|
|
return key
|
|
raise WorkflowError(f"Model did not return requested stem {stem!r}. Available stems: {list(results)}")
|
|
|
|
|
|
def _save_dirs(step: WorkflowStep, stem: str) -> list[str]:
|
|
if not step.save:
|
|
return []
|
|
target = _case_insensitive_get(step.save, stem)
|
|
if target in (None, False, ""):
|
|
return []
|
|
if target is True:
|
|
return [step.id]
|
|
if isinstance(target, list):
|
|
return [str(item).strip() for item in target if str(item).strip()]
|
|
target = str(target).strip()
|
|
return [target] if target else []
|
|
|
|
|
|
def _case_insensitive_get(mapping: dict[str, Any], key: str) -> Any:
|
|
if key in mapping:
|
|
return mapping[key]
|
|
lower = str(key).lower()
|
|
for item_key, value in mapping.items():
|
|
if str(item_key).lower() == lower:
|
|
return value
|
|
return None
|
|
|
|
|
|
def _to_model_audio(audio: Any) -> np.ndarray:
|
|
array = np.asarray(audio, dtype=np.float32)
|
|
if array.ndim == 1:
|
|
return np.ascontiguousarray(array)
|
|
if array.ndim != 2:
|
|
raise WorkflowError(f"Expected mono or stereo audio, got shape {array.shape}.")
|
|
if array.shape[0] in (1, 2):
|
|
return np.ascontiguousarray(array)
|
|
if array.shape[1] in (1, 2):
|
|
return np.ascontiguousarray(array.T)
|
|
raise WorkflowError(f"Expected mono or stereo audio, got shape {array.shape}.")
|
|
|
|
|
|
def _to_save_audio(audio: Any) -> np.ndarray:
|
|
array = np.asarray(audio, dtype=np.float32)
|
|
if array.ndim == 1:
|
|
return np.ascontiguousarray(array)
|
|
if array.ndim != 2:
|
|
raise WorkflowError(f"Expected mono or stereo audio, got shape {array.shape}.")
|
|
if array.shape[1] in (1, 2):
|
|
return np.ascontiguousarray(array)
|
|
if array.shape[0] in (1, 2):
|
|
return np.ascontiguousarray(array.T)
|
|
raise WorkflowError(f"Expected mono or stereo audio, got shape {array.shape}.")
|
|
|
|
|
|
def _ensure_sample_rate(audio: np.ndarray, current_sr: int, target_sr: int) -> np.ndarray:
|
|
if int(current_sr) == int(target_sr):
|
|
return audio
|
|
import librosa
|
|
|
|
return np.ascontiguousarray(
|
|
librosa.resample(np.asarray(audio, dtype=np.float32), orig_sr=int(current_sr), target_sr=int(target_sr), axis=-1)
|
|
)
|
|
|
|
|
|
def _input_files(input_path: str | os.PathLike) -> list[str]:
|
|
path = Path(input_path)
|
|
if path.is_file():
|
|
return [str(path)]
|
|
if path.is_dir():
|
|
return [str(item) for item in sorted(path.iterdir()) if item.is_file()]
|
|
raise WorkflowError(f"Input path does not exist: {path}")
|
|
|
|
|
|
def _unique_track_names(paths: list[str]) -> list[str]:
|
|
original_stems = {Path(path).stem for path in paths}
|
|
next_suffix: dict[str, int] = {}
|
|
used: set[str] = set()
|
|
names = []
|
|
for path in paths:
|
|
stem = Path(path).stem
|
|
if stem not in used:
|
|
used.add(stem)
|
|
names.append(stem)
|
|
continue
|
|
suffix = next_suffix.get(stem, 2)
|
|
candidate = f"{stem}_{suffix}"
|
|
while candidate in used or candidate in original_stems:
|
|
suffix += 1
|
|
candidate = f"{stem}_{suffix}"
|
|
next_suffix[stem] = suffix + 1
|
|
used.add(candidate)
|
|
names.append(candidate)
|
|
return names
|
|
|
|
|
|
def _step_option(workflow: Workflow, step: WorkflowStep, key: str, override: Any = None) -> Any:
|
|
value = getattr(step, key, None)
|
|
if value is not None:
|
|
return value
|
|
if override is not None:
|
|
return override
|
|
return workflow.defaults.get(key)
|
|
|
|
|
|
def _merged_inference_params(workflow: Workflow, step: WorkflowStep) -> dict[str, Any]:
|
|
defaults = workflow.defaults.get("inference_params") or {}
|
|
return {**defaults, **(step.inference_params or {})}
|
|
|
|
|
|
def _safe_filename_part(value: str) -> str:
|
|
safe = re.sub(r"[\\/:\0]+", "_", str(value)).strip()
|
|
return safe or "stem"
|
|
|
|
|
|
def _default_separator_factory(model_name: str, **kwargs):
|
|
model_type = kwargs.pop("model_type", None)
|
|
model_path = kwargs.pop("model_path", None)
|
|
config_path = kwargs.pop("config_path", None)
|
|
if model_path:
|
|
kwargs.pop("model_dir", None)
|
|
from .separator import MSSeparator
|
|
|
|
return MSSeparator(model_type=model_type, model_path=model_path, config_path=config_path, **kwargs)
|
|
|
|
from .model_registry import create_separator
|
|
return create_separator(model_name, **kwargs)
|
|
|
|
|
|
def _default_model_resolver(*args, **kwargs):
|
|
from .model_registry import resolve_model
|
|
|
|
return resolve_model(*args, **kwargs)
|
|
|
|
|
|
def _default_audio_loader(*args, **kwargs):
|
|
from .audio_io import load_audio
|
|
|
|
return load_audio(*args, **kwargs)
|
|
|
|
|
|
def _default_audio_saver(*args, **kwargs):
|
|
from .audio_io import save_audio
|
|
|
|
return save_audio(*args, **kwargs)
|