Spaces:
Sleeping
Sleeping
| import io | |
| import os | |
| import numpy as np | |
| import torch | |
| import librosa | |
| import gradio as gr | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| from huggingface_hub import hf_hub_download | |
| from gradio_client import utils as grc_utils | |
| # IMPORTANT: | |
| # Your Space needs access to VocosVocoderModule. | |
| # Options: | |
| # 1) Vendor minimal code into the Space (recommended for reproducibility), or | |
| # 2) pip install your GitHub repo in requirements.txt (see below). | |
| from src.modules import VocosVocoderModule | |
| MODEL_REPO_ID = os.getenv("MODEL_REPO_ID", "warisqr007/StreamingVocos") | |
| CKPT_FILENAME = os.getenv("CKPT_FILENAME", "epoch=3.ckpt") | |
| SAMPLE_RATE = int(os.getenv("SAMPLE_RATE", "16000")) | |
| # def safe_get_type(schema): | |
| # # First handle booleans (previous error) | |
| # if isinstance(schema, bool): | |
| # return "boolean" | |
| # if "const" in schema: | |
| # return "const" | |
| # if "enum" in schema: | |
| # return "enum" | |
| # if "type" in schema: | |
| # return schema["type"] | |
| # if schema.get("$ref"): | |
| # return "$ref" | |
| # if schema.get("oneOf"): | |
| # return "oneOf" | |
| # if schema.get("anyOf"): | |
| # return "anyOf" | |
| # if schema.get("allOf"): | |
| # return "allOf" | |
| # if "type" not in schema: | |
| # return {} | |
| # raise grc_utils.APIInfoParseError(f"Cannot parse type for {schema}") | |
| # grc_utils.get_type = safe_get_type | |
| # orig__json_schema_to_python_type = grc_utils._json_schema_to_python_type | |
| # def patched__json_schema_to_python_type(schema, defs=None): | |
| # # Handle the specific "string or null" pattern that caused your crash | |
| # if isinstance(schema, dict) and schema.get("anyOf"): | |
| # # Extract non-null types | |
| # non_null = [s for s in schema["anyOf"] if s.get("type") != "null"] | |
| # if len(non_null) == 1 and non_null[0].get("type") == "string": | |
| # # Represent as optional string | |
| # return "str | None" | |
| # # Fallback to original behavior if it's some other pattern | |
| # return orig__json_schema_to_python_type(schema, defs) | |
| # grc_utils._json_schema_to_python_type = patched__json_schema_to_python_type | |
| # def _fig_to_rgb_array(fig): | |
| # """Convert a matplotlib figure to an RGB numpy array.""" | |
| # fig.canvas.draw() | |
| # w, h = fig.canvas.get_width_height() | |
| # img = np.frombuffer(fig.canvas.tostring_rgb(), dtype=np.uint8).reshape(h, w, 3) | |
| # return img | |
| def _fig_to_rgb_array(fig): | |
| """Convert a matplotlib figure to an RGB numpy array (robust across matplotlib versions).""" | |
| fig.canvas.draw() | |
| # Newer matplotlib: use RGBA buffer | |
| if hasattr(fig.canvas, "buffer_rgba"): | |
| buf = np.asarray(fig.canvas.buffer_rgba()) # (H, W, 4) RGBA | |
| return buf[..., :3].copy() # (H, W, 3) RGB | |
| # Fallback: some versions expose ARGB | |
| if hasattr(fig.canvas, "tostring_argb"): | |
| w, h = fig.canvas.get_width_height() | |
| argb = np.frombuffer(fig.canvas.tostring_argb(), dtype=np.uint8).reshape(h, w, 4) | |
| # ARGB -> RGBA | |
| rgba = argb[:, :, [1, 2, 3, 0]] | |
| return rgba[..., :3].copy() | |
| raise RuntimeError("Unsupported matplotlib canvas: cannot extract pixel buffer.") | |
| def get_model(): | |
| if hasattr(get_model, "_model") and get_model._model is not None: | |
| return get_model._model | |
| try: | |
| ckpt_path = hf_hub_download(repo_id=MODEL_REPO_ID, filename=CKPT_FILENAME) | |
| model = VocosVocoderModule.load_from_checkpoint(ckpt_path, map_location="cpu") | |
| model.eval() | |
| get_model._model = model | |
| return model | |
| except Exception as e: | |
| raise RuntimeError(f"Model load failed: {e}") | |
| def mel_to_image(mel_80_t): | |
| """ | |
| mel_80_t: torch.Tensor shaped (80, T) or numpy shaped (80, T) | |
| returns: RGB numpy image | |
| """ | |
| if isinstance(mel_80_t, torch.Tensor): | |
| mel_np = mel_80_t.detach().cpu().numpy() | |
| else: | |
| mel_np = mel_80_t | |
| fig = plt.figure(figsize=(8, 3)) | |
| plt.imshow(mel_np, aspect="auto", origin="lower") | |
| plt.xlabel("Time frames") | |
| plt.ylabel("Mel bins") | |
| plt.title("Mel-spectrogram") | |
| plt.colorbar() | |
| plt.tight_layout() | |
| img = _fig_to_rgb_array(fig) | |
| plt.close(fig) | |
| return img | |
| def compute_mel(model, audio_1d_np): | |
| """ | |
| audio_1d_np: float numpy array shape (T,) | |
| returns mel_spec: torch.Tensor shape (1, 80, Tm) or whatever feature_extractor returns | |
| """ | |
| audio_t = torch.from_numpy(audio_1d_np).float().unsqueeze(0).unsqueeze(0) # (1,1,T) | |
| mel = model.feature_extractor(audio_t) # expected (1,80,Tm) or similar | |
| return mel | |
| def run_reconstruct(audio_path, chunk_size): | |
| """ | |
| Gradio callback. | |
| Returns: | |
| input_audio (sr, np) | |
| input_mel_img | |
| output_audio (sr, np) | |
| output_mel_img | |
| """ | |
| if audio_path is None or str(audio_path).strip() == "": | |
| return None, None, None, None | |
| try: | |
| model = get_model() | |
| except Exception as e: | |
| return None, np.zeros((10,10,3), dtype=np.uint8), None, np.zeros((10,10,3), dtype=np.uint8) | |
| # Load input audio | |
| x, _ = librosa.load(audio_path, sr=SAMPLE_RATE, mono=True) | |
| # Input mel | |
| mel_in = compute_mel(model, x) | |
| # Make (80, T) for plotting | |
| mel_in_80t = mel_in.squeeze(0) # (80, Tm) if mel is (1,80,Tm) | |
| in_mel_img = mel_to_image(mel_in_80t) | |
| # Reconstruct (streaming chunk mode) | |
| chunk_size = int(chunk_size) | |
| if chunk_size < 1: | |
| chunk_size = 1 | |
| # Your notebook uses model.decoder[0] and model.decoder[1] streaming contexts. | |
| # Keep identical behavior here. | |
| y_chunks = [] | |
| with model.decoder[0].streaming(batch_size=1), model.decoder[1].streaming(batch_size=1): | |
| for mel_chunk in mel_in.split(chunk_size, dim=2): | |
| y_chunks.append(model(mel_chunk)) | |
| y = torch.cat(y_chunks, dim=2).squeeze().cpu().numpy() | |
| # Output mel (computed from reconstructed audio) | |
| mel_out = compute_mel(model, y) | |
| mel_out_80t = mel_out.squeeze(0) | |
| out_mel_img = mel_to_image(mel_out_80t) | |
| return (SAMPLE_RATE, x), in_mel_img, (SAMPLE_RATE, y), out_mel_img | |
| with gr.Blocks() as demo: | |
| gr.Markdown( | |
| """ | |
| # 🎙️ Streaming Vocos (Demo) | |
| Upload or record audio, then click **Reconstruct**. | |
| **Left:** input waveform + input mel | |
| **Right:** reconstructed waveform + reconstructed mel | |
| """ | |
| ) | |
| with gr.Row(): | |
| with gr.Column(): | |
| in_audio = gr.Audio( | |
| sources=["upload", "microphone"], | |
| type="filepath", | |
| label="Input audio (upload or record)", | |
| ) | |
| in_mel = gr.Image(label="Input mel-spectrogram", type="numpy") | |
| with gr.Column(): | |
| out_audio = gr.Audio(label="Reconstructed audio", type="numpy") | |
| out_mel = gr.Image(label="Reconstructed mel-spectrogram", type="numpy") | |
| with gr.Row(): | |
| chunk = gr.Slider( | |
| minimum=1, | |
| maximum=50, | |
| value=1, | |
| step=1, | |
| label="Streaming chunk size (mel frames per chunk)", | |
| ) | |
| btn = gr.Button("Reconstruct", variant="primary") | |
| btn.click( | |
| fn=run_reconstruct, | |
| inputs=[in_audio, chunk], | |
| outputs=[in_audio, in_mel, out_audio, out_mel], | |
| ) | |
| demo.queue() | |
| demo.launch(server_name="0.0.0.0", server_port=int(os.environ.get("PORT", 7860))) | |