import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") import spaces # noqa: E402 (must precede torch) import torch # noqa: E402 import io # noqa: E402 import re # noqa: E402 import time # noqa: E402 import gradio as gr # noqa: E402 import librosa # noqa: E402 import matplotlib # noqa: E402 matplotlib.use("Agg") import matplotlib.pyplot as plt # noqa: E402 import numpy as np # noqa: E402 from PIL import Image # noqa: E402 from peft import PeftModel # noqa: E402 from spotsound import ( # noqa: E402 DETECTION_PROMPT, GROUNDING_PROMPT, AudioFlamingo3ForTemporalConditionalGeneration, AudioFlamingo3TemporalProcessor, build_conversation, ) BASE_MODEL = "nvidia/audio-flamingo-3-hf" ADAPTER = "Loie/SpotSound" SR = 16000 DEFAULT_MAX_SECONDS = 300 GROUNDING = "Temporal grounding (when?)" DETECTION = "Event detection (does it occur?)" # ---------------------------------------------------------------- model ------ processor = AudioFlamingo3TemporalProcessor.from_pretrained(BASE_MODEL) model = AudioFlamingo3ForTemporalConditionalGeneration.from_pretrained( BASE_MODEL, dtype=torch.bfloat16, attn_implementation="sdpa" ) # `torch_device="cpu"`: peft otherwise infers "cuda" and safetensors would try to # materialise the adapter on a GPU that does not exist yet under ZeroGPU. model = PeftModel.from_pretrained(model, ADAPTER, torch_device="cpu").merge_and_unload() model = model.eval().to("cuda") # ------------------------------------------------------------- utilities ----- _INTERVAL_RE = re.compile( r"from\s*(-?\d+(?:\.\d+)?)\s*s?\s*(?:econds)?\s*to\s*(-?\d+(?:\.\d+)?)\s*s?", re.I ) _PAIR_RE = re.compile(r"(-?\d+(?:\.\d+)?)\s*s\s*(?:-|to|–)\s*(-?\d+(?:\.\d+)?)\s*s", re.I) def parse_intervals(answer: str, duration: float): """Extract `[start, end]` second-pairs from SpotSound's textual answer.""" matches = _INTERVAL_RE.findall(answer) or _PAIR_RE.findall(answer) intervals = [] for start, end in matches: s, e = float(start), float(end) if e < s: s, e = e, s s = max(0.0, min(s, duration)) e = max(0.0, min(e, duration)) if e - s > 1e-3: intervals.append((s, e)) return intervals def plot_waveform(wav: np.ndarray, duration: float, intervals, query: str) -> Image.Image: """Render the waveform with the predicted temporal windows highlighted.""" n_bins = 1800 step = max(1, len(wav) // n_bins) trimmed = wav[: step * (len(wav) // step)] env = np.abs(trimmed.reshape(-1, step)).max(axis=1) if len(trimmed) else np.zeros(1) env = env / (env.max() + 1e-8) t = np.linspace(0, duration, len(env)) fig, ax = plt.subplots(figsize=(12, 3.1), dpi=110) ax.fill_between(t, -env, env, color="#b9bfcc", linewidth=0) for start, end in intervals: ax.axvspan(start, end, color="#f97316", alpha=0.30, linewidth=0) for x in (start, end): ax.axvline(x, color="#ea580c", linewidth=1.3) ax.text( (start + end) / 2, 1.12, f"{start:.2f}s – {end:.2f}s", ha="center", va="bottom", fontsize=9, color="#9a3412", fontweight="bold", ) title = f'"{query}"' if query else "query" ax.set_title( f"SpotSound — {title}" + (f" · {len(intervals)} window(s) found" if intervals else " · no window predicted"), fontsize=11, ) ax.set_xlim(0, max(duration, 1e-3)) ax.set_ylim(-1.35, 1.35) ax.set_yticks([]) ax.set_xlabel("time (seconds)") for side in ("top", "right", "left"): ax.spines[side].set_visible(False) fig.tight_layout() buf = io.BytesIO() fig.savefig(buf, format="png") plt.close(fig) buf.seek(0) return Image.open(buf).convert("RGB") def extract_segments(wav: np.ndarray, intervals): """Concatenate the predicted windows (0.25 s of silence between them).""" if not intervals: return wav gap = np.zeros(int(0.25 * SR), dtype=np.float32) pieces = [] for start, end in intervals: piece = wav[int(start * SR) : int(end * SR)] if len(piece): pieces.append(piece) pieces.append(gap) if not pieces: return wav return np.concatenate(pieces[:-1]).astype(np.float32) def _estimate_duration( audio_path, query="", task=GROUNDING, max_audio_seconds=DEFAULT_MAX_SECONDS, max_new_tokens=128, *args, **kwargs, ): seconds = 60.0 try: seconds = min(float(librosa.get_duration(path=audio_path)), float(max_audio_seconds)) except Exception: pass # Measured on ZeroGPU (128 new tokens): 1.2 s @ 18 s audio, 1.4 s @ 90 s, 1.9 s @ 300 s; # the rest of the budget covers decode length and pre/post-processing. return int(min(90, 8 + 0.05 * seconds + 0.05 * float(max_new_tokens))) # ------------------------------------------------------------- inference ----- @spaces.GPU(duration=_estimate_duration) def spot( audio_path: str, query: str, task: str = GROUNDING, max_audio_seconds: int = DEFAULT_MAX_SECONDS, max_new_tokens: int = 128, progress=gr.Progress(track_tqdm=True), ): """Localise a sound event described in natural language inside an audio recording. Args: audio_path: path to the audio file to search through. query: natural-language description of the sound to look for, e.g. "dog barking". task: "Temporal grounding (when?)" to get timestamps, or "Event detection (does it occur?)" for a yes/no answer. max_audio_seconds: audio longer than this is truncated before inference. max_new_tokens: generation budget for the answer. Returns: A waveform image with the predicted windows highlighted, the raw model answer, and the audio cropped to the predicted windows. """ if audio_path is None: raise gr.Error("Please provide an audio file.") query = (query or "").strip() if not query: raise gr.Error("Please describe the sound you are looking for.") wav, _ = librosa.load(audio_path, sr=SR, mono=True) wav = np.asarray(wav, dtype=np.float32) max_samples = int(max(1, int(max_audio_seconds)) * SR) truncated = len(wav) > max_samples wav = wav[:max_samples] duration = len(wav) / SR prompt = DETECTION_PROMPT if task == DETECTION else GROUNDING_PROMPT conversation = build_conversation(wav, query, prompt=prompt) inputs = processor.apply_chat_template( conversation, tokenize=True, add_generation_prompt=True, return_dict=True ).to("cuda").to(model.dtype) started = time.perf_counter() with torch.inference_mode(): outputs = model.generate( **inputs, max_new_tokens=int(max_new_tokens), do_sample=False ) answer = processor.batch_decode( outputs[:, inputs["input_ids"].shape[1] :], skip_special_tokens=True )[0].strip() elapsed = time.perf_counter() - started intervals = parse_intervals(answer, duration) if task != DETECTION else [] plot = plot_waveform(wav, duration, intervals, query) lines = [f"Answer: {answer}"] if intervals: lines.append( "Windows: " + ", ".join(f"[{s:.2f}s → {e:.2f}s]" for s, e in intervals) ) elif task != DETECTION: lines.append("No temporal window could be parsed from the answer.") lines.append( f"Audio: {duration:.1f}s{' (truncated)' if truncated else ''} · " f"inference: {elapsed:.1f}s" ) report = "\n".join(lines) return plot, report, (SR, extract_segments(wav, intervals)) # --------------------------------------------------------------------- UI ---- CSS = """ #col-container { max-width: 1180px; margin: 0 auto; } .dark .gradio-container { color: var(--body-text-color); } """ with gr.Blocks() as demo: with gr.Column(elem_id="col-container"): gr.Markdown( """ # 🔍🔊 SpotSound — fine-grained audio temporal grounding Find *when* a described sound happens inside a long recording. SpotSound is a LoRA on top of **Audio Flamingo 3** that interleaves explicit timestamps into the audio stream, so the model can answer with precise start/end times for short events buried in dense background noise. [Paper](https://huggingface.co/papers/2604.13023) · [Project page](https://loiesun.github.io/spotsound/) · [Code](https://github.com/LoieSun/SpotSound) · [Model](https://huggingface.co/Loie/SpotSound) · [Benchmark](https://huggingface.co/datasets/Loie/SpotSound-Bench) """ ) with gr.Row(): with gr.Column(scale=1): audio_in = gr.Audio( label="Audio recording", type="filepath", sources=["upload", "microphone"] ) query_in = gr.Textbox( label="Sound to look for", placeholder="e.g. dog barking, police car siren, hair dryer drying…", lines=1, ) run_btn = gr.Button("Spot it", variant="primary") with gr.Accordion("Advanced settings", open=False): task_in = gr.Radio( choices=[GROUNDING, DETECTION], value=GROUNDING, label="Task", ) max_seconds_in = gr.Slider( 30, 600, value=DEFAULT_MAX_SECONDS, step=30, label="Truncate audio to (seconds)", ) max_tokens_in = gr.Slider( 16, 512, value=128, step=16, label="Max new tokens" ) with gr.Column(scale=1): plot_out = gr.Image(label="Predicted temporal window(s)", type="pil") answer_out = gr.Textbox(label="Model answer", lines=4) segment_out = gr.Audio( label="Spotted segment(s) — full clip if nothing was found", type="numpy", ) gr.Examples( examples=[ ["examples/_Uro9suV3xU_130_187.wav", "hair dryer drying"], ["examples/ClTzzGQatXo_30_48.wav", "ambulance siren"], ["examples/eXQYEfqCU08_38_68.wav", "train horning"], ["examples/fTcSVQJ2h8g_0_90.wav", "police car siren"], ], inputs=[audio_in, query_in], outputs=[plot_out, answer_out, segment_out], fn=spot, cache_examples=True, cache_mode="lazy", label="Examples from SpotSound-Bench", ) inputs = [audio_in, query_in, task_in, max_seconds_in, max_tokens_in] outputs = [plot_out, answer_out, segment_out] run_btn.click(fn=spot, inputs=inputs, outputs=outputs, api_name="spot") query_in.submit(fn=spot, inputs=inputs, outputs=outputs, api_name=False) demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)