codec / app.py
humair025's picture
Update app.py
6f3c0ac verified
Raw
History Blame Contribute Delete
11.4 kB
"""
BlueCodec — speech autoencoder demo
Upload audio, encode it into BlueCodec's 24-dimensional ~86 Hz latent
space, decode it back to a waveform, and inspect what happened in
between: the reconstructed audio, the waveform comparison, the latent
representation itself, and the numbers behind the compression.
"""
# `spaces` must be imported before torch (or anything that touches CUDA) so
# its ZeroGPU shims are installed first. This Space runs on ZeroGPU hardware,
# which requires at least one @spaces.GPU-decorated function.
import spaces
import os
import time
import tempfile
import urllib.request
import numpy as np
import torch
import torchaudio
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from matplotlib.colors import LinearSegmentedColormap
import gradio as gr
from bluecodec import BlueCodec
# --------------------------------------------------------------------------
# Config
# --------------------------------------------------------------------------
MODEL_ID = "notmax123/blue-codec"
TARGET_SR = 44100
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
EXAMPLE_DIR = "examples_cache"
EXAMPLE_URL = (
"https://github.com/thewh1teagle/phonikud-chatterbox/releases/"
"download/asset-files-v1/female1.wav"
)
EXAMPLE_PATH = os.path.join(EXAMPLE_DIR, "female1.wav")
NAVY = "#0B0F1A"
SURFACE = "#111827"
BLUE = "#3B9EFF"
VIOLET = "#8B7CFF"
TEXT = "#E6EBF5"
MUTED = "#7C89A3"
LINE = "#22304A"
# --------------------------------------------------------------------------
# Model loading (once, at import time)
# --------------------------------------------------------------------------
_codec = None
_load_error = None
try:
_codec = BlueCodec.from_pretrained(MODEL_ID, device=DEVICE)
except Exception as exc: # noqa: BLE001 - surface any load failure in the UI
_load_error = str(exc)
def get_codec():
return _codec, _load_error
def ensure_example():
"""Best-effort download of a sample clip so the Space ships with an example."""
try:
os.makedirs(EXAMPLE_DIR, exist_ok=True)
if not os.path.exists(EXAMPLE_PATH):
urllib.request.urlretrieve(EXAMPLE_URL, EXAMPLE_PATH)
return EXAMPLE_PATH
except Exception: # noqa: BLE001 - examples are optional, never block startup
return None
# --------------------------------------------------------------------------
# Plotting helpers
# --------------------------------------------------------------------------
def _style_ax(ax):
ax.set_facecolor(SURFACE)
for spine in ax.spines.values():
spine.set_color(LINE)
spine.set_linewidth(0.8)
ax.tick_params(colors=MUTED, labelsize=8)
ax.xaxis.label.set_color(MUTED)
ax.yaxis.label.set_color(MUTED)
ax.grid(color=LINE, linewidth=0.5, alpha=0.6)
ax.set_axisbelow(True)
def make_waveform_figure(original: np.ndarray, reconstructed: np.ndarray, sr: int):
t_o = np.arange(len(original)) / sr
t_r = np.arange(len(reconstructed)) / sr
fig, axes = plt.subplots(2, 1, figsize=(9, 4), sharex=True, facecolor=NAVY)
axes[0].plot(t_o, original, color=BLUE, linewidth=0.6)
axes[0].set_ylabel("Original")
axes[1].plot(t_r, reconstructed, color=VIOLET, linewidth=0.6)
axes[1].set_ylabel("Reconstructed")
axes[1].set_xlabel("Time (s)")
for ax in axes:
_style_ax(ax)
ax.set_ylim(-1.05, 1.05)
fig.tight_layout()
plt.close(fig)
return fig
def latent_to_matrix(latents: torch.Tensor) -> np.ndarray:
"""Return latents as a (24, T) numpy array regardless of axis order."""
arr = latents.detach().float().cpu().numpy()
arr = np.squeeze(arr)
if arr.ndim != 2:
raise ValueError(f"Unexpected latent shape: {tuple(latents.shape)}")
if arr.shape[0] > arr.shape[1]:
arr = arr.T
return arr
def make_latent_figure(latent_matrix: np.ndarray):
cmap = LinearSegmentedColormap.from_list(
"bluecodec", [NAVY, BLUE, VIOLET, "#F2F5FF"]
)
fig, ax = plt.subplots(figsize=(9, 2.8), facecolor=NAVY)
im = ax.imshow(latent_matrix, aspect="auto", cmap=cmap, interpolation="nearest")
ax.set_ylabel("Latent dim")
ax.set_xlabel("Frame (~86 Hz)")
_style_ax(ax)
ax.grid(False)
cbar = fig.colorbar(im, ax=ax, fraction=0.025, pad=0.02)
cbar.ax.tick_params(colors=MUTED, labelsize=7)
cbar.outline.set_edgecolor(LINE)
fig.tight_layout()
plt.close(fig)
return fig
# --------------------------------------------------------------------------
# Core inference
# --------------------------------------------------------------------------
@spaces.GPU(duration=60)
def process(audio_path, progress=gr.Progress(track_tqdm=False)):
if audio_path is None:
raise gr.Error("Upload or record a clip first.")
codec, err = get_codec()
if codec is None:
raise gr.Error(f"BlueCodec failed to load on startup: {err}")
progress(0.05, desc="Loading audio")
wav, sr = torchaudio.load(audio_path, backend="soundfile")
channels_in = wav.shape[0]
if channels_in > 1:
wav = wav.mean(dim=0, keepdim=True)
if sr != TARGET_SR:
wav = torchaudio.functional.resample(wav, sr, TARGET_SR)
wav = wav.to(DEVICE)
duration_s = wav.shape[-1] / TARGET_SR
progress(0.25, desc="Encoding to latents")
t0 = time.perf_counter()
with torch.inference_mode():
latents = codec.encode(wav)
t1 = time.perf_counter()
progress(0.6, desc="Decoding back to audio")
with torch.inference_mode():
recon = codec.decode(latents)
t2 = time.perf_counter()
encode_s, decode_s = t1 - t0, t2 - t1
total_s = t2 - t0
progress(0.9, desc="Rendering plots")
recon_cpu = recon.detach().float().cpu().clamp(-1, 1)
out_path = tempfile.NamedTemporaryFile(suffix=".wav", delete=False).name
torchaudio.save(out_path, recon_cpu, TARGET_SR)
orig_np = wav.detach().float().cpu().numpy()[0]
recon_np = recon_cpu.numpy()[0]
n_samples = orig_np.shape[-1]
original_bits = n_samples * 16 # reference: 16-bit PCM, mono
latent_bits = latents.numel() * 32 # latents assumed float32
ratio = original_bits / latent_bits if latent_bits else float("nan")
rtf = duration_s / total_s if total_s > 0 else float("nan")
stats_md = f"""
| | |
|---|---|
| Duration | {duration_s:.2f} s |
| Channels | mono ({"downmixed from " + str(channels_in) if channels_in > 1 else "native"}) |
| Device | {DEVICE.upper()} |
| Latent shape | {tuple(latents.shape)} |
| Latent rate | ~86 Hz × 24 dims |
| Encode time | {encode_s * 1000:.0f} ms |
| Decode time | {decode_s * 1000:.0f} ms |
| Compression ratio | {ratio:.1f}× vs 16-bit PCM |
| Real-time factor | {rtf:.1f}× |
"""
wave_fig = make_waveform_figure(orig_np, recon_np, TARGET_SR)
latent_fig = make_latent_figure(latent_to_matrix(latents))
return out_path, stats_md, wave_fig, latent_fig
# --------------------------------------------------------------------------
# UI
# --------------------------------------------------------------------------
CUSTOM_CSS = f"""
@import url('https://fonts.googleapis.com/css2?family=Space+Grotesk:wght@500;600;700&family=IBM+Plex+Sans:wght@400;500&display=swap');
.gradio-container {{
background: {NAVY} !important;
color: {TEXT} !important;
font-family: 'IBM Plex Sans', sans-serif !important;
}}
#bc-hero h1 {{
font-family: 'Space Grotesk', sans-serif !important;
font-weight: 700 !important;
font-size: 2.15rem !important;
letter-spacing: -0.01em;
color: {TEXT} !important;
margin-bottom: 0.2rem !important;
}}
#bc-hero p {{
color: {MUTED} !important;
font-size: 1.02rem;
max-width: 62ch;
}}
#bc-section-label h3 {{
color: {MUTED} !important;
font-family: 'Space Grotesk', sans-serif !important;
font-weight: 500 !important;
font-size: 0.95rem !important;
margin: 0.6rem 0 0.3rem 0 !important;
}}
.bc-panel {{
background: {SURFACE} !important;
border: 1px solid {LINE} !important;
border-radius: 6px !important;
box-shadow: none !important;
}}
.bc-stats table {{
color: {TEXT} !important;
font-family: 'Space Grotesk', sans-serif !important;
width: 100%;
border-collapse: collapse;
}}
.bc-stats table, .bc-stats th, .bc-stats td {{
border: none !important;
}}
.bc-stats tr {{
border-bottom: 1px solid {LINE} !important;
}}
.bc-stats td:first-child {{
color: {MUTED} !important;
font-size: 0.82rem;
font-family: 'IBM Plex Sans', sans-serif !important;
}}
.bc-stats td:last-child {{
text-align: right;
font-variant-numeric: tabular-nums;
}}
#bc-cta {{
background: {BLUE} !important;
border: none !important;
color: #051220 !important;
font-weight: 600 !important;
}}
#bc-cta:hover {{
background: {VIOLET} !important;
}}
#bc-footer, #bc-footer p {{
color: {MUTED} !important;
font-size: 0.85rem;
}}
#bc-footer a {{
color: {BLUE} !important;
}}
"""
with gr.Blocks(title="BlueCodec — speech autoencoder") as demo:
gr.Markdown(
"# BlueCodec\n"
"Upload a voice clip and pass it through a 24-dimensional latent "
"bottleneck at roughly 86 frames per second, then back out as audio. "
"Everything in between — the reconstructed sound, the waveforms, and "
"the latent representation itself — is rendered below.",
elem_id="bc-hero",
)
if _load_error:
gr.Markdown(
f"**Model failed to load at startup:** `{_load_error}`\n\n"
"The interface below will report this error again if you try to run it.",
)
with gr.Row():
with gr.Column(scale=1):
audio_in = gr.Audio(
sources=["upload", "microphone"],
type="filepath",
label="Original audio",
elem_classes=["bc-panel"],
)
run_btn = gr.Button("Encode → decode", variant="primary", elem_id="bc-cta")
example_path = ensure_example()
if example_path:
gr.Examples(
examples=[[example_path]],
inputs=[audio_in],
label="Try a sample clip",
)
with gr.Column(scale=1):
audio_out = gr.Audio(label="Reconstructed audio", elem_classes=["bc-panel"])
stats_out = gr.Markdown(elem_classes=["bc-panel", "bc-stats"])
gr.Markdown("### Waveform: original vs. reconstructed", elem_id="bc-section-label")
wave_plot = gr.Plot(elem_classes=["bc-panel"])
gr.Markdown("### Latent representation (24 dims × ~86 Hz)", elem_id="bc-section-label")
latent_plot = gr.Plot(elem_classes=["bc-panel"])
gr.Markdown(
"This Space runs the BlueCodec autoencoder. Model weights are hosted at "
"[notmax123/blue-codec](https://huggingface.co/notmax123/blue-codec) on "
"Hugging Face, and the source lives at "
"[maxmelichov/blue-codec](https://github.com/maxmelichov/blue-codec) on GitHub. "
"The compression ratio shown is measured against 16-bit PCM at the same "
"sample rate, and assumes float32 latents.",
elem_id="bc-footer",
)
run_btn.click(
fn=process,
inputs=[audio_in],
outputs=[audio_out, stats_out, wave_plot, latent_plot],
)
if __name__ == "__main__":
demo.queue().launch(theme=gr.themes.Base(), css=CUSTOM_CSS)