from __future__ import annotations import base64 import html import io import json import os import re import tempfile import time from dataclasses import asdict, dataclass from pathlib import Path from typing import Any import av import gradio as gr import numpy as np import soundfile as sf import spaces import torch from transformers import AutoModelForCausalLM, AutoProcessor MODEL_ID = os.getenv("MODEL_ID", "OpenMOSS-Team/MOSS-Transcribe-Diarize") SAMPLE_RATE = 16_000 MAX_INPUT_SECONDS = int(os.getenv("MAX_INPUT_SECONDS", "900")) # 15 min demo guardrail DEFAULT_MAX_NEW_TOKENS = int(os.getenv("MAX_NEW_TOKENS", "8192")) # This is the model's documented default diarization prompt. DEFAULT_PROMPT = ( "请将音频转写为文本,每一段需以起始时间戳和说话人编号" "([S01]、[S02]、[S03]…)开头,正文为对应的语音内容," "并在段末标注结束时间戳,以清晰标明该段语音范围。" ) SPEAKER_COLORS = [ "#2563EB", "#D97706", "#059669", "#DC2626", "#7C3AED", "#0891B2", "#DB2777", "#4F46E5", ] @dataclass(frozen=True) class Segment: start: float end: float speaker: str text: str @property def duration(self) -> float: return max(0.0, self.end - self.start) # ZeroGPU recommends placing the model on CUDA at module scope. Outside a GPU # function, Spaces provides CUDA emulation so this is safe on ZeroGPU. DTYPE = torch.bfloat16 PROCESSOR = AutoProcessor.from_pretrained(MODEL_ID, trust_remote_code=True) MODEL = ( AutoModelForCausalLM.from_pretrained( MODEL_ID, trust_remote_code=True, dtype="auto", ) .to(dtype=DTYPE) .to("cuda") .eval() ) def _media_path(value: Any) -> str: """Normalize Gradio File values across Gradio versions.""" if value is None: raise gr.Error("Upload an audio or video file first.") if isinstance(value, str): return value if isinstance(value, dict): candidate = value.get("path") or value.get("name") if candidate: return str(candidate) candidate = getattr(value, "name", None) if candidate: return str(candidate) return str(value) def _decode_audio(path: str, sample_rate: int = SAMPLE_RATE) -> np.ndarray: """Decode the first audio stream to mono float32 using PyAV.""" chunks: list[np.ndarray] = [] with av.open(path) as container: stream = next((s for s in container.streams if s.type == "audio"), None) if stream is None: raise gr.Error("The uploaded file does not contain an audio stream.") resampler = av.audio.resampler.AudioResampler( format="fltp", layout="mono", rate=sample_rate, ) for frame in container.decode(stream): converted = resampler.resample(frame) if converted is None: continue if not isinstance(converted, list): converted = [converted] for item in converted: chunks.append(item.to_ndarray().reshape(-1).astype(np.float32, copy=False)) tail = resampler.resample(None) if tail is not None: if not isinstance(tail, list): tail = [tail] for item in tail: chunks.append(item.to_ndarray().reshape(-1).astype(np.float32, copy=False)) if not chunks: raise gr.Error("No decodable audio samples were found.") audio = np.concatenate(chunks).astype(np.float32, copy=False) peak = float(np.max(np.abs(audio))) if audio.size else 0.0 # Some decoders can return float planar audio outside [-1, 1]. Keep playback sane. if peak > 1.25: audio = audio / peak return audio def _write_wav(audio: np.ndarray, sample_rate: int = SAMPLE_RATE, prefix: str = "moss") -> str: fd, path = tempfile.mkstemp(prefix=f"{prefix}-", suffix=".wav") os.close(fd) sf.write(path, audio, sample_rate, subtype="PCM_16") return path def _format_time(seconds: float) -> str: seconds = max(0.0, float(seconds)) minutes, sec = divmod(seconds, 60) hours, minutes = divmod(int(minutes), 60) if hours: return f"{hours:02d}:{minutes:02d}:{sec:05.2f}" return f"{minutes:02d}:{sec:05.2f}" def _parse_transcript(text: str) -> list[Segment]: """Parse canonical `[start][Sxx]text[end]` MOSS output.""" pattern = re.compile( r"\[(?P\d+(?:\.\d+)?)\]\s*" r"\[(?PS\d+)\]\s*" r"(?P.*?)" r"\[(?P\d+(?:\.\d+)?)\]", flags=re.IGNORECASE | re.DOTALL, ) segments: list[Segment] = [] for match in pattern.finditer(text): start = float(match.group("start")) end = float(match.group("end")) if end <= start: continue body = re.sub(r"\s+", " ", match.group("text")).strip() segments.append( Segment( start=start, end=end, speaker=match.group("speaker").upper(), text=body, ) ) return segments def _build_messages(audio_path: str, prompt: str) -> list[dict[str, Any]]: return [ { "role": "user", "content": [ {"type": "audio", "audio": audio_path}, {"type": "text", "text": prompt}, ], } ] def _speaker_map(segments: list[Segment]) -> dict[str, str]: speakers = list(dict.fromkeys(segment.speaker for segment in segments)) return {speaker: SPEAKER_COLORS[i % len(SPEAKER_COLORS)] for i, speaker in enumerate(speakers)} def _timeline_html(segments: list[Segment], audio_duration: float) -> str: if not segments: return "
No parsed speaker turns.
" colors = _speaker_map(segments) duration = max(audio_duration, max(segment.end for segment in segments), 0.001) speakers = list(colors) lanes: list[str] = [] for speaker in speakers: bars: list[str] = [] for index, segment in enumerate(segments, start=1): if segment.speaker != speaker: continue left = min(100.0, max(0.0, (segment.start / duration) * 100.0)) width = max(0.35, min(100.0 - left, (segment.duration / duration) * 100.0)) title = html.escape( f"#{index} {speaker} · {_format_time(segment.start)}–{_format_time(segment.end)} · {segment.text}", quote=True, ) bars.append( f"" ) lanes.append( "
" f"
{html.escape(speaker)}
" f"
{''.join(bars)}
" "
" ) return ( "
" "
Speaker timelineHover a block for turn details
" f"{''.join(lanes)}" "
0:0025%50%75%" f"{html.escape(_format_time(duration))}
" "
" ) def _transcript_html(segments: list[Segment]) -> str: if not segments: return "
The raw model output is available below, but no canonical speaker turns were parsed.
" colors = _speaker_map(segments) cards: list[str] = [] previous_speaker: str | None = None for index, segment in enumerate(segments, start=1): speaker_changed = segment.speaker != previous_speaker cards.append( "
" "
" f"{html.escape(segment.speaker)}" f"#{index:02d}" f"{html.escape(_format_time(segment.start))} → {html.escape(_format_time(segment.end))}" f"{segment.duration:.2f}s" "
" f"
{html.escape(segment.text) if segment.text else '(no text)'}
" "
" ) previous_speaker = segment.speaker if speaker_changed else previous_speaker return f"
{''.join(cards)}
" def _summary_html(segments: list[Segment], duration: float, elapsed: float, token_count: int) -> str: speakers = len({segment.speaker for segment in segments}) speech_seconds = sum(segment.duration for segment in segments) rtf = elapsed / duration if duration > 0 else 0.0 return ( "
" f"
Speakers{speakers}
" f"
Turns{len(segments)}
" f"
Audio{html.escape(_format_time(duration))}
" f"
Tagged speech{speech_seconds:.1f}s
" f"
Inference{elapsed:.1f}s
" f"
RTF{rtf:.2f}×
" f"
Generated{token_count} tok
" "
" ) def _turn_choices(segments: list[Segment]) -> list[tuple[str, int]]: choices: list[tuple[str, int]] = [] for index, segment in enumerate(segments): preview = segment.text[:72] + ("…" if len(segment.text) > 72 else "") label = f"#{index + 1:02d} · {segment.speaker} · {_format_time(segment.start)}–{_format_time(segment.end)} · {preview}" choices.append((label, index)) return choices def _result_json_file(segments: list[Segment], raw_text: str) -> str: fd, path = tempfile.mkstemp(prefix="moss-diarization-", suffix=".json") os.close(fd) payload = { "model": MODEL_ID, "segments": [asdict(segment) | {"duration": segment.duration} for segment in segments], "raw_text": raw_text, } Path(path).write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8") return path def _gpu_duration( media: Any, max_new_tokens: float, hotwords: str, progress: gr.Progress = gr.Progress(), ) -> int: """Conservative dynamic ZeroGPU reservation, capped for a public demo.""" try: path = _media_path(media) with av.open(path) as container: seconds = float(container.duration / av.time_base) if container.duration else 60.0 except Exception: seconds = 60.0 token_factor = max(1.0, float(max_new_tokens or DEFAULT_MAX_NEW_TOKENS) / 8192.0) return int(min(300, max(45, 35 + seconds * 0.65 * token_factor))) @spaces.GPU(duration=_gpu_duration) def transcribe(media: Any, max_new_tokens: float, hotwords: str, progress: gr.Progress = gr.Progress()): source_path = _media_path(media) progress(0.05, desc="Decoding audio") audio = _decode_audio(source_path, SAMPLE_RATE) duration = len(audio) / SAMPLE_RATE if duration > MAX_INPUT_SECONDS: raise gr.Error( f"This demo is capped at {MAX_INPUT_SECONDS // 60} minutes per upload to keep ZeroGPU usage predictable. " "Raise MAX_INPUT_SECONDS if you control the Space and want longer files." ) playback_path = _write_wav(audio, SAMPLE_RATE, prefix="moss-source") prompt = DEFAULT_PROMPT if hotwords and hotwords.strip(): prompt += f"热词提示:{hotwords.strip()}" progress(0.15, desc="Preparing MOSS inputs") messages = _build_messages(playback_path, prompt) rendered = PROCESSOR.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) inputs = PROCESSOR( text=rendered, audio=[audio], max_length=131072, audio_kwargs={"device": "cuda"}, return_tensors="pt", ).to("cuda") prompt_len = int(inputs["attention_mask"][0].sum().item()) progress(0.25, desc="Running MOSS-Transcribe-Diarize") started = time.perf_counter() with torch.inference_mode(), torch.amp.autocast("cuda", dtype=DTYPE): output_ids = MODEL.generate( input_ids=inputs["input_ids"], attention_mask=inputs["attention_mask"], input_features=inputs["input_features"], audio_feature_lengths=inputs["audio_feature_lengths"], audio_chunk_mapping=inputs["audio_chunk_mapping"], max_new_tokens=max(256, int(max_new_tokens or DEFAULT_MAX_NEW_TOKENS)), do_sample=False, ) elapsed = time.perf_counter() - started generated = output_ids[0][prompt_len:] raw_text = PROCESSOR.tokenizer.decode(generated, skip_special_tokens=True).strip() segments = _parse_transcript(raw_text) progress(0.92, desc="Building diarization review") state = { "audio_path": playback_path, "sample_rate": SAMPLE_RATE, "segments": [asdict(segment) for segment in segments], } choices = _turn_choices(segments) first_value = 0 if choices else None status = ( f"**Done.** Parsed **{len(segments)}** diarized turns across " f"**{len({segment.speaker for segment in segments})}** speakers." if segments else "**Done, but the output did not match the canonical timestamp/speaker format.** Check Raw model output." ) return ( state, playback_path, _summary_html(segments, duration, elapsed, int(generated.numel())), _timeline_html(segments, duration), _transcript_html(segments), gr.update(choices=choices, value=first_value, interactive=bool(choices)), raw_text, _result_json_file(segments, raw_text), status, ) def _restore_segments(state: dict[str, Any] | None) -> list[Segment]: if not state: return [] return [Segment(**item) for item in state.get("segments", [])] def show_turn(index: int | float | None, state: dict[str, Any] | None): segments = _restore_segments(state) if index is None or not segments or not state: return None, "Select a turn to inspect it in isolation." idx = min(max(int(index), 0), len(segments) - 1) segment = segments[idx] audio, file_sr = sf.read(str(state["audio_path"]), dtype="float32", always_2d=False) if audio.ndim > 1: audio = np.mean(audio, axis=1, dtype=np.float32) sr = int(file_sr or state.get("sample_rate", SAMPLE_RATE)) start = min(len(audio), max(0, round(segment.start * sr))) end = min(len(audio), max(start, round(segment.end * sr))) clip_path = _write_wav(audio[start:end], sr, prefix=f"turn-{idx + 1:02d}-{segment.speaker}") detail = ( f"**Turn {idx + 1}/{len(segments)} · {segment.speaker}** \n" f"{_format_time(segment.start)} → {_format_time(segment.end)} · {segment.duration:.2f}s \n" f"{segment.text or '_(no text)_'}" ) return clip_path, detail def step_turn(index: int | float | None, state: dict[str, Any] | None, offset: int): segments = _restore_segments(state) if not segments: return gr.update() current = int(index or 0) target = min(max(current + offset, 0), len(segments) - 1) return gr.update(value=target) CSS = """ .gradio-container {max-width: 1180px !important;} .hero {padding: 8px 0 2px;} .hero h1 {font-size: 2rem; margin-bottom: .2rem;} .hero p {color: var(--body-text-color-subdued); margin-top: 0;} .stats-grid {display:grid;grid-template-columns:repeat(auto-fit,minmax(120px,1fr));gap:10px;margin:10px 0 18px;} .stat {border:1px solid var(--border-color-primary);border-radius:12px;padding:12px 13px;background:var(--background-fill-primary);} .stat span {display:block;font-size:.74rem;color:var(--body-text-color-subdued);text-transform:uppercase;letter-spacing:.04em;} .stat strong {display:block;font-size:1.12rem;margin-top:3px;} .timeline-card {border:1px solid var(--border-color-primary);border-radius:14px;padding:14px;background:var(--background-fill-primary);margin-bottom:14px;} .timeline-head {display:flex;justify-content:space-between;gap:16px;align-items:center;margin-bottom:10px;} .timeline-head span {font-size:.78rem;color:var(--body-text-color-subdued);} .tl-row {display:grid;grid-template-columns:72px 1fr;gap:8px;align-items:center;margin:7px 0;} .tl-label {font-weight:650;font-size:.82rem;display:flex;align-items:center;gap:6px;} .speaker-dot {width:9px;height:9px;border-radius:50%;display:inline-block;} .tl-track {height:24px;position:relative;border-radius:7px;background:var(--background-fill-secondary);overflow:hidden;} .tl-bar {position:absolute;top:3px;height:18px;border-radius:5px;min-width:3px;opacity:.92;} .tl-axis {margin-left:80px;display:flex;justify-content:space-between;color:var(--body-text-color-subdued);font-size:.68rem;} .turn-list {display:flex;flex-direction:column;gap:8px;} .turn-card {border:1px solid var(--border-color-primary);border-radius:12px;padding:11px 13px;background:var(--background-fill-primary);} .turn-meta {display:flex;gap:7px;align-items:center;flex-wrap:wrap;margin-bottom:6px;} .speaker-pill {color:#fff;font-weight:700;font-size:.75rem;padding:3px 8px;border-radius:999px;} .turn-number,.turn-time,.turn-duration {font-size:.76rem;color:var(--body-text-color-subdued);} .turn-duration {margin-left:auto;} .turn-text {font-size:.96rem;line-height:1.5;} .empty-card {border:1px dashed var(--border-color-primary);border-radius:12px;padding:18px;color:var(--body-text-color-subdued);} #source-audio audio, #turn-audio audio {min-height:54px;} """ def build_app() -> gr.Blocks: with gr.Blocks(title="MOSS Diarization Review", css=CSS) as demo: result_state = gr.State({}) gr.HTML( "

MOSS Diarization Review

" "

Upload a conversation, run MOSS-Transcribe-Diarize on ZeroGPU, and inspect who spoke when.

" ) with gr.Row(equal_height=False): with gr.Column(scale=5): media = gr.File( label="Audio or video", file_types=["audio", "video"], type="filepath", ) with gr.Row(): run_button = gr.Button("Transcribe & diarize", variant="primary", scale=4) clear_button = gr.ClearButton(value="Clear", components=[media], scale=1) with gr.Accordion("Advanced", open=False): hotwords = gr.Textbox( label="Hotwords / names (optional)", placeholder="OpenMOSS, Hugging Face, product names…", ) max_new_tokens = gr.Slider( 1024, 16384, value=DEFAULT_MAX_NEW_TOKENS, step=1024, label="Max new tokens", info="Increase for long recordings with many turns.", ) with gr.Column(scale=4): gr.Markdown( "**Review flow** \n" "1. Use the timeline to spot speaker switches. \n" "2. Scan the color-coded turn list. \n" "3. Select any turn to listen to that slice alone." ) status = gr.Markdown() source_audio = gr.Audio(label="Source audio", type="filepath", elem_id="source-audio") stats = gr.HTML() gr.Markdown("## Diarization") timeline = gr.HTML() with gr.Row(): previous_turn = gr.Button("← Previous", scale=1) turn_selector = gr.Dropdown( label="Inspect a single turn", choices=[], interactive=False, scale=6, ) next_turn = gr.Button("Next →", scale=1) with gr.Row(equal_height=False): turn_audio = gr.Audio(label="Selected turn", type="filepath", elem_id="turn-audio") turn_detail = gr.Markdown("Select a turn to inspect it in isolation.") gr.Markdown("## Speaker-aware transcript") transcript = gr.HTML() with gr.Accordion("Raw model output / export", open=False): raw_output = gr.Textbox(label="Raw MOSS output", lines=8) json_file = gr.File(label="Diarization JSON") clear_button.add([ result_state, source_audio, stats, timeline, transcript, turn_selector, turn_audio, turn_detail, raw_output, json_file, status ]) outputs = [ result_state, source_audio, stats, timeline, transcript, turn_selector, raw_output, json_file, status, ] run_event = run_button.click( transcribe, inputs=[media, max_new_tokens, hotwords], outputs=outputs, ) run_event.then(show_turn, inputs=[turn_selector, result_state], outputs=[turn_audio, turn_detail]) turn_selector.input(show_turn, inputs=[turn_selector, result_state], outputs=[turn_audio, turn_detail]) previous_turn.click( lambda index, state: step_turn(index, state, -1), inputs=[turn_selector, result_state], outputs=turn_selector, ).then(show_turn, inputs=[turn_selector, result_state], outputs=[turn_audio, turn_detail]) next_turn.click( lambda index, state: step_turn(index, state, 1), inputs=[turn_selector, result_state], outputs=turn_selector, ).then(show_turn, inputs=[turn_selector, result_state], outputs=[turn_audio, turn_detail]) return demo if __name__ == "__main__": build_app().queue(default_concurrency_limit=1).launch()