import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") import spaces # noqa: E402 (must precede torch / CUDA-touching imports) import json # noqa: E402 import re # noqa: E402 import shutil # noqa: E402 import threading # noqa: E402 import time # noqa: E402 import gradio as gr # noqa: E402 import librosa # noqa: E402 import torch # noqa: E402 from transformers import ( # noqa: E402 AutoProcessor, Qwen3OmniMoeForConditionalGeneration, TextIteratorStreamer, ) # -------------------------------------------------------------------------------------- # Model # -------------------------------------------------------------------------------------- HUB_ID = "PleasedPenguin/A2R-30B-A3B" _total, _used, _free = shutil.disk_usage("/") print(f"[boot] disk: total={_total / 2**30:.1f}GB used={_used / 2**30:.1f}GB free={_free / 2**30:.1f}GB") def _resolve_model_path() -> str: """Fetch the checkpoint onto local disk so `from_pretrained` can mmap the shards. Reading the shards straight off an attached read-only model volume was tried first and does not work: the volume serves truncated small files and OOMs the loader. """ from huggingface_hub import snapshot_download t0 = time.perf_counter() path = snapshot_download( HUB_ID, ignore_patterns=["*.md", "*.png", "*.jpg", "*.gif"], max_workers=8, ) print(f"[boot] snapshot ready in {time.perf_counter() - t0:.1f}s -> {path}") return path MODEL_PATH = _resolve_model_path() print(f"[boot] loading {MODEL_PATH} ...") _t0 = time.perf_counter() processor = AutoProcessor.from_pretrained(MODEL_PATH) model = Qwen3OmniMoeForConditionalGeneration.from_pretrained( MODEL_PATH, dtype=torch.bfloat16, attn_implementation="sdpa", ).eval() model = model.to("cuda") print(f"[boot] model loaded in {time.perf_counter() - _t0:.1f}s (talker: {model.has_talker})") _total, _used, _free = shutil.disk_usage("/") print(f"[boot] disk after load: used={_used / 2**30:.1f}GB free={_free / 2**30:.1f}GB") # Qwen's official Omni system prompt. The paper's evaluation harness # (github.com/dwsmart32/HEAR) applies exactly this to A2R. SYSTEM_PROMPT = ( "You are Qwen, a virtual human developed by the Qwen Team, Alibaba Group, capable of " "perceiving auditory and visual inputs, as well as generating text and speech." ) MAX_AUDIO_SECONDS = 120 TARGET_SR = 16000 # The merged checkpoint ships no generation_config, so `<|im_end|>` has to be supplied # explicitly or the thinker never stops. The authors' harness passes the same id. IM_END_ID = 151645 ENDOFTEXT_ID = 151643 ANSWER_RE = re.compile(r"(.*?)", re.S) REASONING_RE = re.compile(r"(.*?)(?:|$)", re.S) FENCE_RE = re.compile(r"```(?:json)?\s*(.*?)\s*```", re.S) JSON_ANSWER_RE = re.compile(r'"Answer"\s*:\s*"([^"]+)"') def _split_output(text: str): """Split A2R's `...{...}` output. A2R does not always emit the structured wrapper -- for plain multiple-choice prompts it often replies with a bare ```json {"Answer": "B"} ``` block -- so both shapes are handled. """ reasoning_match = REASONING_RE.search(text) reasoning = reasoning_match.group(1).strip() if reasoning_match else "" answer_blocks = ANSWER_RE.findall(text) answer = answer_blocks[-1].strip() if answer_blocks else "" if not answer: fences = FENCE_RE.findall(text) if fences: answer = fences[-1].strip() if not reasoning and not answer: # Nothing structured at all: show the raw completion as the trace. return text.strip(), "" if not reasoning: # Everything before the answer block is the model's own line of thought. head = text.split("```")[0].split("")[0] reasoning = head.replace("", "").replace("", "").strip() if answer: letter = JSON_ANSWER_RE.search(answer) if letter: answer = f"### {letter.group(1)}" else: answer = f"```\n{answer}\n```" if "\n" in answer else f"### {answer}" return reasoning, answer def _estimate_duration(audio, question, max_new_tokens=768, *args, **kwargs): try: n = int(max_new_tokens) except Exception: n = 768 # ~20 tok/s measured on ZeroGPU xlarge, plus ~15 s to stream the packed weights in. return int(min(200, 22 + n * 0.055)) @spaces.GPU(duration=_estimate_duration, size="xlarge") def analyze( audio: str, question: str, max_new_tokens: int = 768, temperature: float = 0.6, top_p: float = 0.95, ): """Answer a speaker-attribution question about a multi-speaker audio clip. Args: audio: path to an audio file containing one or more speakers. question: the question to ask about who is speaking, optionally with multiple-choice options. max_new_tokens: maximum number of tokens A2R may generate. temperature: sampling temperature. top_p: nucleus sampling probability mass. Returns: The model's answer and its speaker-attribution reasoning trace. """ if not audio: raise gr.Error("Please provide an audio clip.") if not question or not question.strip(): raise gr.Error("Please ask a question about the audio.") waveform, _ = librosa.load(audio, sr=TARGET_SR, mono=True) if waveform.shape[0] > MAX_AUDIO_SECONDS * TARGET_SR: waveform = waveform[: MAX_AUDIO_SECONDS * TARGET_SR] gr.Info(f"Audio truncated to the first {MAX_AUDIO_SECONDS} seconds.") messages = [ {"role": "system", "content": SYSTEM_PROMPT}, { "role": "user", "content": [ {"type": "audio", "audio": audio}, {"type": "text", "text": question.strip()}, ], }, ] text = processor.apply_chat_template(messages, add_generation_prompt=True, tokenize=False) inputs = processor( text=[text], audio=[waveform], sampling_rate=TARGET_SR, return_tensors="pt", ).to("cuda") streamer = TextIteratorStreamer( processor.tokenizer, skip_prompt=True, skip_special_tokens=True ) kwargs = dict( **inputs, streamer=streamer, max_new_tokens=int(max_new_tokens), do_sample=float(temperature) > 0, temperature=float(temperature), top_p=float(top_p), eos_token_id=[IM_END_ID, ENDOFTEXT_ID], pad_token_id=ENDOFTEXT_ID, ) thread = threading.Thread(target=model.thinker.generate, kwargs=kwargs, daemon=True) t0 = time.perf_counter() thread.start() acc = "" n_chunks = 0 yield "*Listening…*", "" for chunk in streamer: acc += chunk n_chunks += 1 if n_chunks % 4 == 0: reasoning, answer = _split_output(acc) # For a free-form question A2R answers in prose and there is no separate # answer block: the prose itself is the answer. yield answer or reasoning or "*Reasoning…*", acc thread.join() elapsed = time.perf_counter() - t0 reasoning, answer = _split_output(acc) print(f"[infer] {n_chunks} chunks in {elapsed:.1f}s -> {acc[:120]!r}") yield answer or reasoning or "*(empty response)*", acc # -------------------------------------------------------------------------------------- # UI # -------------------------------------------------------------------------------------- MCQ_TAIL = '\n\nRespond in JSON format, e.g. {"Answer": "A"}' HEAD = "Listen to the main audio and answer the following multiple-choice question." COUNT_Q = ( HEAD + "\n\nHow many distinct speakers are there in the audio?" "\n\nOptions:\n(A) 1\n(B) 2\n(C) 3\n(D) 4" + MCQ_TAIL ) EXAMPLES = [ ["examples/ami_two_speakers.wav", COUNT_Q], ["examples/ami_one_speaker.wav", COUNT_Q], [ "examples/ami_male_first_returns.wav", HEAD + "\n\nIs the first person who speaks in the recording male or female?" "\n\nOptions:\n(A) Male\n(B) Female" + MCQ_TAIL, ], [ "examples/ami_two_speakers.wav", "Walk through this recording turn by turn. How many different voices do you hear, " "and how would you characterise each one? Base your answer on the voices " "themselves, not on what is being said.", ], ] CSS = """ #col-container { max-width: 1100px; margin: 0 auto; } .dark .gradio-container { color: var(--body-text-color); } """ with gr.Blocks(title="A2R — speaker-attributed reasoning") as demo: with gr.Column(elem_id="col-container"): gr.Markdown( """ # A2R-30B-A3B — who said what? Ask questions about **who** is speaking in a multi-party recording, not only about what is said. [A2R](https://huggingface.co/PleasedPenguin/A2R-30B-A3B) is `Qwen3-Omni-30B-A3B-Instruct` trained with GRPO on counterfactual audio with speaker-level hard negatives, so it grounds its answer in vocal cues instead of the transcript. [Paper](https://huggingface.co/papers/2608.29120) · [Project page](https://attributetoreason.github.io/AttributeToReason/) · [Code](https://github.com/dwsmart32/HEAR) · [Model](https://huggingface.co/PleasedPenguin/A2R-30B-A3B) """ ) with gr.Row(): with gr.Column(scale=1): audio_in = gr.Audio( label="Multi-speaker audio", type="filepath", sources=["upload", "microphone"] ) question_in = gr.Textbox( label="Question", lines=7, placeholder=( "e.g. How many different people speak in this recording, and in " "what order do they take turns?" ), ) run = gr.Button("Analyse", variant="primary") with gr.Column(scale=1): gr.Markdown("### Answer") answer_out = gr.Markdown(value="", min_height=90) reasoning_out = gr.Textbox( label="Raw model output", lines=16, buttons=["copy"] ) with gr.Accordion("Advanced settings", open=False): max_new_tokens = gr.Slider( 128, 2048, value=768, step=64, label="Max new tokens" ) temperature = gr.Slider(0.0, 1.5, value=0.6, step=0.05, label="Temperature") top_p = gr.Slider(0.1, 1.0, value=0.95, step=0.01, label="Top-p") gr.Examples( examples=EXAMPLES, inputs=[audio_in, question_in], outputs=[answer_out, reasoning_out], fn=analyze, cache_examples=True, cache_mode="lazy", label="Examples (AMI Meeting Corpus, CC BY 4.0)", ) gr.Markdown( "Example clips are excerpts of the " "[AMI Meeting Corpus](https://huggingface.co/datasets/diarizers-community/ami) " "(CC BY 4.0), one of the source corpora behind the HEAR benchmark. The HEAR " "benchmark audio itself is not redistributable and is therefore not bundled here." ) run.click( analyze, inputs=[audio_in, question_in, max_new_tokens, temperature, top_p], outputs=[answer_out, reasoning_out], api_name="analyze", ) demo.queue(max_size=12).launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)