CoolFace
Apppublic

aiqtech/SoulX-Singer

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
webui.py567 linesDownload Raw Back to root
1import os2import re3import random4import shutil5import sys6import traceback7from pathlib import Path8from typing import Tuple9import spaces10 11import numpy as np12import torch13import librosa14import soundfile as sf15import gradio as gr16 17from preprocess.pipeline import PreprocessPipeline18from soulxsinger.utils.file_utils import load_config19from cli.inference import build_model as build_svs_model, process as svs_process20 21 22ROOT = Path(__file__).parent23 24 25def _get_device() -> str:26    if torch.cuda.is_available():27        return "cuda:0"28    try:29        from spaces.config import Config30        if Config.zero_gpu:31            return "cuda:0"32    except (ImportError, AttributeError):33        pass34    return "cpu"35 36 37def _session_dir_from_target(target_audio_path: str) -> Path:38    stem = Path(target_audio_path).stem39    safe = re.sub(r"[^\w\-]", "_", stem)40    safe = re.sub(r"_+", "_", safe).strip("_") or "session"41    return ROOT / "outputs" / "gradio" / safe[:64]42 43 44class AppState:45    def __init__(self) -> None:46        self.device = _get_device()47        self.preprocess_pipeline = PreprocessPipeline(48            device=self.device,49            language="English",50            save_dir=str(ROOT / "outputs" / "gradio" / "_placeholder" / "transcriptions"),51            vocal_sep=True,52            max_merge_duration=60000,53        )54        config = load_config("soulxsinger/config/soulxsinger.yaml")55        self.svs_config = config56        self.svs_model = build_svs_model(57            model_path="pretrained_models/SoulX-Singer/model.pt",58            config=config,59            device=self.device,60        )61        self.phoneset_path = "soulxsinger/utils/phoneme/phone_set.json"62 63    def run_preprocess(64        self,65        prompt_path: Path,66        target_path: Path,67        session_base: Path,68        prompt_vocal_sep: bool,69        target_vocal_sep: bool,70        prompt_lyric_lang: str,71        target_lyric_lang: str,72    ) -> Tuple[bool, str]:73        try:74            self.preprocess_pipeline.save_dir = str(session_base / "transcriptions" / "prompt")75            self.preprocess_pipeline.run(76                audio_path=str(prompt_path),77                vocal_sep=prompt_vocal_sep,78                max_merge_duration=20000,79                language=prompt_lyric_lang or "English",80            )81            self.preprocess_pipeline.save_dir = str(session_base / "transcriptions" / "target")82            self.preprocess_pipeline.run(83                audio_path=str(target_path),84                vocal_sep=target_vocal_sep,85                max_merge_duration=60000,86                language=target_lyric_lang or "English",87            )88            return True, "preprocess done"89        except Exception as e:90            return False, f"preprocess failed: {e}"91 92    def run_svs(93        self,94        control: str,95        session_base: Path,96        auto_shift: bool,97        pitch_shift: int,98    ) -> Tuple[bool, str, Path | None, Path | None, Path | None]:99        if control not in ("melody", "score"):100            control = "score"101        save_dir = session_base / "generated"102        save_dir.mkdir(parents=True, exist_ok=True)103 104        class Args:105            pass106 107        args = Args()108        args.device = self.device109        args.model_path = "pretrained_models/SoulX-Singer/model.pt"110        args.config = "soulxsinger/config/soulxsinger.yaml"111        args.prompt_wav_path = str(session_base / "audio" / "prompt.wav")112        prompt_meta_path = session_base / "transcriptions" / "prompt" / "metadata.json"113        target_meta_path = session_base / "transcriptions" / "target" / "metadata.json"114        args.prompt_metadata_path = str(prompt_meta_path)115        args.target_metadata_path = str(target_meta_path)116        args.phoneset_path = self.phoneset_path117        args.save_dir = str(save_dir)118        args.auto_shift = auto_shift119        args.pitch_shift = int(pitch_shift)120        args.control = control121        try:122            svs_process(args, self.svs_config, self.svs_model)123            generated = save_dir / "generated.wav"124            if not generated.exists():125                return False, f"inference finished but {generated} not found", None, prompt_meta_path, target_meta_path126            return True, "svs inference done", generated, prompt_meta_path, target_meta_path127        except Exception as e:128            return False, f"svs inference failed: {e}", None, prompt_meta_path, target_meta_path129 130    def run_svs_from_paths(131        self,132        prompt_wav_path: str,133        prompt_metadata_path: str,134        target_metadata_path: str,135        control: str,136        auto_shift: bool,137        pitch_shift: int,138        save_dir: Path | None = None,139    ) -> Tuple[bool, str, Path | None]:140        if save_dir is None:141            import uuid142            save_dir = ROOT / "outputs" / "gradio" / "synthesis" / str(uuid.uuid4())[:8]143        save_dir = Path(save_dir)144        audio_dir = save_dir / "audio"145        prompt_meta_dir = save_dir / "transcriptions" / "prompt"146        target_meta_dir = save_dir / "transcriptions" / "target"147        audio_dir.mkdir(parents=True, exist_ok=True)148        prompt_meta_dir.mkdir(parents=True, exist_ok=True)149        target_meta_dir.mkdir(parents=True, exist_ok=True)150        shutil.copy2(prompt_wav_path, audio_dir / "prompt.wav")151        shutil.copy2(prompt_metadata_path, prompt_meta_dir / "metadata.json")152        shutil.copy2(target_metadata_path, target_meta_dir / "metadata.json")153        ok, msg, merged, _, _ = self.run_svs(154            control=control,155            session_base=save_dir,156            auto_shift=auto_shift,157            pitch_shift=pitch_shift,158        )159        if not ok or merged is None:160            return False, msg or "svs failed", None161        return True, "svs inference done", merged162 163 164from ensure_models import ensure_pretrained_models165ensure_pretrained_models()166 167APP_STATE = AppState()168 169 170def _resolve_file_path(x):171    if x is None:172        return None173    if isinstance(x, tuple):174        x = x[0]175    return x if (x and os.path.isfile(x)) else None176 177 178def _run_transcription_internal(179    prompt_audio, target_audio,180    prompt_lyric_lang, target_lyric_lang,181    prompt_vocal_sep, target_vocal_sep,182):183    """Run transcription, return (prompt_meta_path, target_meta_path) or (None, None)."""184    if isinstance(prompt_audio, tuple):185        prompt_audio = prompt_audio[0]186    if isinstance(target_audio, tuple):187        target_audio = target_audio[0]188 189    session_base = _session_dir_from_target(target_audio)190    audio_dir = session_base / "audio"191    audio_dir.mkdir(parents=True, exist_ok=True)192 193    SR = 44100194    PROMPT_MAX_SEC = 30195    TARGET_MAX_SEC = 60196    prompt_audio_data, _ = librosa.load(prompt_audio, sr=SR, mono=True)197    target_audio_data, _ = librosa.load(target_audio, sr=SR, mono=True)198    prompt_audio_data = prompt_audio_data[: PROMPT_MAX_SEC * SR]199    target_audio_data = target_audio_data[: TARGET_MAX_SEC * SR]200    sf.write(audio_dir / "prompt.wav", prompt_audio_data, SR)201    sf.write(audio_dir / "target.wav", target_audio_data, SR)202 203    ok, msg = APP_STATE.run_preprocess(204        audio_dir / "prompt.wav",205        audio_dir / "target.wav",206        session_base,207        prompt_vocal_sep=prompt_vocal_sep,208        target_vocal_sep=target_vocal_sep,209        prompt_lyric_lang=prompt_lyric_lang or "English",210        target_lyric_lang=target_lyric_lang or "English",211    )212    if not ok:213        print(msg, file=sys.stderr, flush=True)214        return None, None215 216    prompt_meta_path = session_base / "transcriptions" / "prompt" / "metadata.json"217    target_meta_path = session_base / "transcriptions" / "target" / "metadata.json"218    p = str(prompt_meta_path) if prompt_meta_path.exists() else None219    t = str(target_meta_path) if target_meta_path.exists() else None220    return p, t221 222 223@spaces.GPU224def transcription_function(225    prompt_audio, target_audio,226    prompt_metadata, target_metadata,227    prompt_lyric_lang, target_lyric_lang,228    prompt_vocal_sep, target_vocal_sep,229):230    """Step 1: Run transcription only; output (prompt_meta_path, target_meta_path)."""231    try:232        if isinstance(prompt_audio, tuple):233            prompt_audio = prompt_audio[0]234        if isinstance(target_audio, tuple):235            target_audio = target_audio[0]236        if prompt_audio is None or target_audio is None:237            gr.Warning(message="Please upload both prompt audio and target audio")238            return None, None239 240        prompt_meta_resolved = _resolve_file_path(prompt_metadata)241        target_meta_resolved = _resolve_file_path(target_metadata)242        use_input_metadata = prompt_meta_resolved is not None and target_meta_resolved is not None243 244        if use_input_metadata:245            session_base = _session_dir_from_target(target_audio)246            audio_dir = session_base / "audio"247            audio_dir.mkdir(parents=True, exist_ok=True)248            SR = 44100249            prompt_audio_data, _ = librosa.load(prompt_audio, sr=SR, mono=True)250            target_audio_data, _ = librosa.load(target_audio, sr=SR, mono=True)251            prompt_audio_data = prompt_audio_data[: 30 * SR]252            target_audio_data = target_audio_data[: 60 * SR]253            sf.write(audio_dir / "prompt.wav", prompt_audio_data, SR)254            sf.write(audio_dir / "target.wav", target_audio_data, SR)255 256            prompt_meta_path = session_base / "transcriptions" / "prompt" / "metadata.json"257            target_meta_path = session_base / "transcriptions" / "target" / "metadata.json"258            (session_base / "transcriptions" / "prompt").mkdir(parents=True, exist_ok=True)259            (session_base / "transcriptions" / "target").mkdir(parents=True, exist_ok=True)260            shutil.copy2(prompt_meta_resolved, prompt_meta_path)261            shutil.copy2(target_meta_resolved, target_meta_path)262            return str(prompt_meta_path), str(target_meta_path)263        else:264            return _run_transcription_internal(265                prompt_audio, target_audio,266                prompt_lyric_lang, target_lyric_lang,267                prompt_vocal_sep, target_vocal_sep,268            )269    except Exception:270        print(traceback.format_exc(), file=sys.stderr, flush=True)271        return None, None272 273 274@spaces.GPU275def synthesis_function(276    prompt_audio,277    target_audio,278    prompt_metadata=None,279    target_metadata=None,280    control="melody",281    auto_shift=True,282    pitch_shift=0,283    seed=12306,284    prompt_lyric_lang="English",285    target_lyric_lang="English",286    prompt_vocal_sep=True,287    target_vocal_sep=True,288):289    """Single-button: runs transcription first if metadata not provided, then synthesis."""290    try:291        if isinstance(prompt_audio, tuple):292            prompt_audio = prompt_audio[0]293        if isinstance(target_audio, tuple):294            target_audio = target_audio[0]295 296        if not prompt_audio or not os.path.isfile(prompt_audio):297            gr.Warning(message="Please upload both prompt audio and target audio")298            return None, gr.update(), gr.update()299        if not target_audio or not os.path.isfile(target_audio):300            gr.Warning(message="Please upload both prompt audio and target audio")301            return None, gr.update(), gr.update()302 303        prompt_meta_path = _resolve_file_path(prompt_metadata)304        target_meta_path = _resolve_file_path(target_metadata)305 306        # Auto-run transcription if metadata not provided307        if not prompt_meta_path or not target_meta_path:308            p, t = _run_transcription_internal(309                prompt_audio, target_audio,310                prompt_lyric_lang, target_lyric_lang,311                prompt_vocal_sep, target_vocal_sep,312            )313            if not p or not t:314                gr.Warning(message="Transcription failed. Check your audio files.")315                return None, gr.update(), gr.update()316            prompt_meta_path = p317            target_meta_path = t318 319        # Prepare prompt wav320        session_base = _session_dir_from_target(target_audio)321        prompt_wav = session_base / "audio" / "prompt.wav"322        if not prompt_wav.exists():323            audio_dir = session_base / "audio"324            audio_dir.mkdir(parents=True, exist_ok=True)325            SR = 44100326            data, _ = librosa.load(prompt_audio, sr=SR, mono=True)327            data = data[: 30 * SR]328            sf.write(prompt_wav, data, SR)329 330        if control not in ("melody", "score"):331            control = "score"332        seed = int(seed)333        torch.manual_seed(seed)334        np.random.seed(seed)335        random.seed(seed)336 337        ok, msg, merged = APP_STATE.run_svs_from_paths(338            prompt_wav_path=str(prompt_wav),339            prompt_metadata_path=prompt_meta_path,340            target_metadata_path=target_meta_path,341            control=control,342            auto_shift=auto_shift,343            pitch_shift=int(pitch_shift),344        )345        if not ok or merged is None:346            print(msg or "synthesis failed", file=sys.stderr, flush=True)347            return None, gr.update(), gr.update()348 349        # Return generated audio + update metadata displays350        return str(merged), prompt_meta_path, target_meta_path351    except Exception:352        print(traceback.format_exc(), file=sys.stderr, flush=True)353        return None, gr.update(), gr.update()354 355 356 357def render_interface() -> gr.Blocks:358    with gr.Blocks(title="SoulX-Singer", theme=gr.themes.Default()) as page:359        gr.HTML(360            '<div style="'361            'text-align: center; '362            'padding: 1.25rem 0 1.5rem; '363            'margin-bottom: 0.5rem;'364            '">'365            '<div style="'366            'display: inline-block; '367            'font-size: 1.75rem; '368            'font-weight: 700; '369            'letter-spacing: 0.02em; '370            'line-height: 1.3;'371            '">SoulX-Singer</div>'372            '<div style="'373            'width: 80px; '374            'height: 3px; '375            'margin: 1rem auto 0; '376            'background: linear-gradient(90deg, transparent, #6366f1, transparent); '377            'border-radius: 2px;'378            '"></div>'379            '</div>'380        )381 382        with gr.Row(equal_height=False):383            # ── Left column: inputs & controls ──384            with gr.Column(scale=1):385                prompt_audio = gr.Audio(386                    label="Prompt audio (reference voice), max 30s",387                    type="filepath",388                    interactive=True,389                )390                target_audio = gr.Audio(391                    label="Target audio (melody / lyrics source), max 60s",392                    type="filepath",393                    interactive=True,394                )395 396                with gr.Row():397                    control_radio = gr.Radio(398                        choices=["melody", "score"],399                        value="melody",400                        label="Control type",401                        scale=1,402                    )403                    auto_shift = gr.Checkbox(404                        label="Auto pitch shift",405                        value=True,406                        interactive=True,407                        scale=1,408                    )409 410                synthesis_btn = gr.Button(411                    value="🎤 Generate singing voice",412                    variant="primary",413                    size="lg",414                )415 416                # ── Advanced: transcription settings & metadata ──417                with gr.Accordion("Advanced: Transcription & Metadata", open=False):418                    with gr.Row():419                        pitch_shift = gr.Number(420                            label="Pitch shift (semitones)",421                            value=0,422                            minimum=-36,423                            maximum=36,424                            step=1,425                            interactive=True,426                            scale=1,427                        )428                        seed_input = gr.Number(429                            label="Seed",430                            value=12306,431                            step=1,432                            interactive=True,433                            scale=1,434                        )435                    gr.Markdown(436                        "Upload your own metadata files to skip automatic transcription. "437                        "You can use the [SoulX-Singer-Midi-Editor]"438                        "(https://huggingface.co/spaces/Soul-AILab/SoulX-Singer-Midi-Editor) "439                        "to edit metadata for better alignment."440                    )441                    with gr.Row():442                        prompt_lyric_lang = gr.Dropdown(443                            label="Prompt lyric language",444                            choices=[445                                ("Mandarin", "Mandarin"),446                                ("Cantonese", "Cantonese"),447                                ("English", "English"),448                            ],449                            value="English",450                            interactive=True,451                            scale=1,452                        )453                        target_lyric_lang = gr.Dropdown(454                            label="Target lyric language",455                            choices=[456                                ("Mandarin", "Mandarin"),457                                ("Cantonese", "Cantonese"),458                                ("English", "English"),459                            ],460                            value="English",461                            interactive=True,462                            scale=1,463                        )464                    with gr.Row():465                        prompt_vocal_sep = gr.Checkbox(466                            label="Prompt vocal separation",467                            value=False,468                            interactive=True,469                            scale=1,470                        )471                        target_vocal_sep = gr.Checkbox(472                            label="Target vocal separation",473                            value=True,474                            interactive=True,475                            scale=1,476                        )477                    transcription_btn = gr.Button(478                        value="Run singing transcription",479                        variant="secondary",480                        size="lg",481                    )482                    with gr.Row():483                        prompt_metadata = gr.File(484                            label="Prompt metadata",485                            type="filepath",486                            file_types=[".json"],487                            interactive=True,488                        )489                        target_metadata = gr.File(490                            label="Target metadata",491                            type="filepath",492                            file_types=[".json"],493                            interactive=True,494                        )495 496            # ── Right column: output ──497            with gr.Column(scale=1):498                output_audio = gr.Audio(499                    label="Generated audio",500                    type="filepath",501                    interactive=False,502                )503                gr.Examples(504                    examples=[505                        ["raven.wav", "happy_birthday.mp3"],506                        ["anita.wav", "happy_birthday.mp3"],507                        ["obama.wav", "happy_birthday.mp3"],508                        ["raven.wav", "everybody_loves.wav"],509                        ["anita.wav", "everybody_loves.wav"],510                        ["obama.wav", "everybody_loves.wav"],511                    ],512                    inputs=[prompt_audio, target_audio],513                    outputs=[output_audio, prompt_metadata, target_metadata],514                    fn=synthesis_function,515                    cache_examples=True,516                    cache_mode="lazy"517                )518 519        # ── Event handlers ──520        prompt_audio.change(521            fn=lambda: None,522            inputs=[],523            outputs=[prompt_metadata],524        )525        526        target_audio.change(527            fn=lambda: None,528            inputs=[],529            outputs=[target_metadata],530        )531        532        transcription_btn.click(533            fn=transcription_function,534            inputs=[535                prompt_audio, target_audio,536                prompt_metadata, target_metadata,537                prompt_lyric_lang, target_lyric_lang,538                prompt_vocal_sep, target_vocal_sep,539            ],540            outputs=[prompt_metadata, target_metadata],541        )542 543        synthesis_btn.click(544            fn=synthesis_function,545            inputs=[546                prompt_audio, target_audio,547                prompt_metadata, target_metadata,548                control_radio, auto_shift, pitch_shift, seed_input,549                prompt_lyric_lang, target_lyric_lang,550                prompt_vocal_sep, target_vocal_sep,551            ],552            outputs=[output_audio, prompt_metadata, target_metadata],553        )554 555    return page556 557 558if __name__ == "__main__":559    import argparse560    parser = argparse.ArgumentParser()561    parser.add_argument("--port", type=int, default=7860, help="Gradio server port")562    parser.add_argument("--share", action="store_true", help="Create public link")563    args = parser.parse_args()564 565    page = render_interface()566    page.queue()567    page.launch(share=args.share, server_name="0.0.0.0", server_port=args.port)