File size: 11,099 Bytes
989544b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bb16670
 
 
989544b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
70b8677
989544b
 
 
 
 
 
 
 
70b8677
 
 
989544b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e9726a4
989544b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e9726a4
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
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)