Spaces:
Running on Zero
Running on Zero
| import sys | |
| sys.stdout.reconfigure(line_buffering=True) | |
| try: | |
| import spaces | |
| except ImportError: | |
| # keep @spaces.GPU usable as a no-op; ZeroGPU requires this exact name. | |
| class spaces: | |
| class GPU: | |
| def __init__(self, func=None, duration=60): | |
| self.func = func | |
| def __call__(self, *args, **kwargs): | |
| if self.func is not None: | |
| return self.func(*args, **kwargs) | |
| func = args[0] | |
| return func | |
| import os | |
| import threading | |
| from types import SimpleNamespace | |
| import torch | |
| import torchaudio | |
| import gradio as gr | |
| from audiotools import AudioSignal | |
| from pyharp import ModelCard, build_endpoint, load_audio, save_audio | |
| from core.models.ebase import EndToEndLightningSystem | |
| from core.models.e2e.bandit.bandit import PasstFiLMConditionedBandit | |
| DEVICE = "cuda" if torch.cuda.is_available() else "cpu" | |
| CKPT_PATH = os.path.join(os.path.dirname(__file__), "ev-pre-aug.ckpt") | |
| MODEL_FS = 44100 | |
| QUERY_LENGTH_SECONDS = 10.0 | |
| # chunked_inference pads the mixture based on chunk/hop size; pad_mode=reflect | |
| # rejects audio below 33 seconds outright. so upfront validation is important. | |
| MIN_MIXTURE_SECONDS = 34.0 | |
| MIN_MIXTURE_SAMPLES = int(MIN_MIXTURE_SECONDS * MODEL_FS) | |
| # inference chunking (default: chunk_size_seconds=6.0, hop_size_seconds=0.5, | |
| # batch_size=12, per repo config config/data/moisesdb-test.yml), an internal | |
| # windowing detail, not something a musician can meaningfully tune. | |
| CHUNK_SIZE_SECONDS = 6.0 | |
| HOP_SIZE_SECONDS = 0.5 | |
| INFERENCE_BATCH_SIZE = 12 | |
| # architecture kwargs, per repo config config/models/bandit-query-pre.yml | |
| MODEL_KWARGS = dict( | |
| in_channel=2, | |
| band_type="musical", | |
| n_bands=64, | |
| additive_film=True, | |
| multiplicative_film=True, | |
| film_depth=2, | |
| n_sqm_modules=8, | |
| emb_dim=128, | |
| rnn_dim=256, | |
| bidirectional=True, | |
| rnn_type="GRU", | |
| mlp_dim=512, | |
| hidden_activation="Tanh", | |
| hidden_activation_kwargs=None, | |
| complex_mask=True, | |
| use_freq_weights=True, | |
| n_fft=2048, | |
| win_length=2048, | |
| hop_length=512, | |
| window_fn="hann_window", | |
| wkwargs=None, | |
| power=None, | |
| center=True, | |
| normalized=True, | |
| pad_mode="reflect", | |
| onesided=True, | |
| fs=MODEL_FS, | |
| pretrain_encoder=None, # repo points this to a cluster for warm-start during training; | |
| # strict=True in load_model overwrites all of these anyway. | |
| freeze_encoder=False, | |
| ) | |
| system = None | |
| model_ready = False # has the model been moved onto the GPU yet? | |
| model_loading = True | |
| model_error = None | |
| def load_model(): | |
| """Builds the model and loads the checkpoint on CPU. The move to GPU happens | |
| later, inside the @spaces.GPU-decorated process_fn, the only place ZeroGPU | |
| intercepts CUDA calls.""" | |
| global system, model_loading, model_error | |
| try: | |
| model = PasstFiLMConditionedBandit(**MODEL_KWARGS) | |
| # actual repo builds real loss/metric/augmentation handlers here | |
| # (train.py::inference_byoq); we only use chunked_inference, so args are | |
| # left as None rather than porting the training-only classes that build them. | |
| system = EndToEndLightningSystem.load_from_checkpoint( | |
| CKPT_PATH, | |
| map_location="cpu", | |
| strict=True, | |
| model=model, | |
| loss_handler=None, | |
| metrics=None, | |
| augmentation_handler=None, | |
| inference_handler=SimpleNamespace( | |
| fs=MODEL_FS, | |
| chunk_size_seconds=CHUNK_SIZE_SECONDS, | |
| hop_size_seconds=HOP_SIZE_SECONDS, | |
| batch_size=INFERENCE_BATCH_SIZE, | |
| ), | |
| optimization_bundle=None, | |
| ) | |
| print("Model loaded (CPU).") | |
| except Exception as e: | |
| model_error = str(e) | |
| print(f"Load error: {e}") | |
| finally: | |
| model_loading = False | |
| threading.Thread(target=load_model, daemon=True).start() | |
| model_card = ModelCard( | |
| name="Banquet", | |
| description=( | |
| "Extracts any instrument from a music mixture using a short audio " | |
| "example as a query, instead of a fixed vocals/drums/bass/other setup. " | |
| f"Mixture must be at least {MIN_MIXTURE_SECONDS:.0f} seconds long, and " | |
| "query audio should be ~10 seconds of the instrument you want extracted." | |
| ), | |
| author="Karn N. Watcharasupat and Alexander Lerch", | |
| tags=["source separation", "music"], | |
| ) | |
| def _load_resampled(path: str) -> torch.Tensor: | |
| """Loads an audio file and resamples it to the model's sample rate. | |
| Returns a (channels, samples) float32 tensor.""" | |
| signal = load_audio(path) | |
| audio = signal.audio_data.squeeze(0) | |
| if signal.sample_rate != MODEL_FS: | |
| audio = torchaudio.functional.resample( | |
| audio, orig_freq=signal.sample_rate, new_freq=MODEL_FS | |
| ) | |
| return audio | |
| def _ensure_stereo(audio: torch.Tensor) -> torch.Tensor: | |
| """The model's architecture is built for a fixed 2-channel input; duplicate | |
| mono uploads to stereo and drop any channels beyond the first two.""" | |
| if audio.shape[0] == 1: | |
| audio = audio.repeat(2, 1) | |
| elif audio.shape[0] > 2: | |
| audio = audio[:2] | |
| return audio | |
| def _fit_query_length(query: torch.Tensor) -> torch.Tensor: | |
| """Truncates or tiles the query to exactly 10 seconds.""" | |
| target_len = int(QUERY_LENGTH_SECONDS * MODEL_FS) | |
| if query.shape[-1] > target_len: | |
| query = query[:, :target_len] | |
| elif query.shape[-1] < target_len: | |
| reps = target_len // query.shape[-1] + 1 | |
| query = query.repeat(1, reps)[:, :target_len] | |
| return query | |
| def process_fn(mixture_path: str, query_path: str) -> str: | |
| """Separates the instrument described by the query clip out of the mixture.""" | |
| global model_ready | |
| if model_loading: | |
| raise gr.Error("Model is still loading, please wait a moment and try again.") | |
| if system is None: | |
| raise gr.Error(f"Model failed to load: {model_error}") | |
| if not model_ready: | |
| system.to(DEVICE) # only safe here, inside @spaces.GPU | |
| model_ready = True | |
| orig_fs = load_audio(mixture_path).sample_rate | |
| mixture = _ensure_stereo(_load_resampled(mixture_path)) | |
| if mixture.shape[-1] < MIN_MIXTURE_SAMPLES: | |
| raise gr.Error( | |
| f"Mixture is too short ({mixture.shape[-1] / MODEL_FS:.1f}s). " | |
| f"Needs to be at least {MIN_MIXTURE_SECONDS:.0f}s long." | |
| ) | |
| mixture = mixture.unsqueeze(0).to(DEVICE) | |
| query = _fit_query_length(_load_resampled(query_path)).unsqueeze(0).to(DEVICE) | |
| batch = { | |
| "mixture": {"audio": mixture}, | |
| "query": {"audio": query}, | |
| "metadata": {"stem": ["target"]}, | |
| "estimates": {}, | |
| } | |
| out = system.chunked_inference(batch) | |
| estimate = out["estimates"]["target"]["audio"].squeeze(0).cpu() | |
| if orig_fs != MODEL_FS: | |
| estimate = torchaudio.functional.resample( | |
| estimate, orig_freq=MODEL_FS, new_freq=orig_fs | |
| ) | |
| output_signal = AudioSignal(estimate, sample_rate=orig_fs) | |
| return save_audio(output_signal) | |
| with gr.Blocks() as demo: | |
| input_components = [ | |
| gr.Audio(type="filepath", label="Mixture").harp_required(True), | |
| gr.Audio( | |
| type="filepath", | |
| label="Query Example (~10s clip of the instrument you want extracted)", | |
| ).harp_required(True), | |
| ] | |
| output_components = [ | |
| gr.Audio( | |
| type="filepath", | |
| label=f"Separated Audio (at least {MIN_MIXTURE_SECONDS:.0f} seconds audio required)" | |
| ).set_info( | |
| "The instrument extracted from the mixture, matched to the query example. " | |
| ), | |
| ] | |
| build_endpoint( | |
| model_card=model_card, | |
| input_components=input_components, | |
| output_components=output_components, | |
| process_fn=process_fn, | |
| ) | |
| demo.queue().launch(pwa=True) | |