ABAB-Annotation / app.py
HighnessAnnotation's picture
Update app.py
b544084 verified
Raw
History Blame
7.51 kB
import tempfile
import traceback
from pathlib import Path
import gradio as gr
from pipeline import get_pipeline, to_json, to_csv, to_abab_text
UPLOAD_DIR = Path(tempfile.gettempdir()) / "speech_annotation"
UPLOAD_DIR.mkdir(parents=True, exist_ok=True)
_last_segments = []
SPEAKER_COLORS = [
("#4F46E5", "#EEF2FF"), # A - indigo
("#059669", "#ECFDF5"), # B - green
("#DC2626", "#FEF2F2"), # C - red
("#D97706", "#FFFBEB"), # D - amber
("#7C3AED", "#F5F3FF"), # E - purple
("#0891B2", "#ECFEFF"), # F - cyan
("#DB2777", "#FDF2F8"), # G - pink
("#65A30D", "#F7FEE7"), # H - lime
("#EA580C", "#FFF7ED"), # I - orange
("#0284C7", "#F0F9FF"), # J - sky
]
def make_conversation_html(segments):
if not segments:
return ""
speaker_list = list(dict.fromkeys(s.speaker for s in segments))
color_map = {spk: SPEAKER_COLORS[i % len(SPEAKER_COLORS)] for i, spk in enumerate(speaker_list)}
# Legend
legend_items = "".join(
f"<span style='display:inline-flex;align-items:center;gap:6px;margin-right:16px'>"
f"<span style='width:12px;height:12px;border-radius:50%;background:{color_map[spk][0]}'></span>"
f"<span style='font-weight:600;color:{color_map[spk][0]}'>Speaker {spk}</span></span>"
for spk in speaker_list
)
legend = f"<div style='padding:12px 16px;border-bottom:1px solid #e5e7eb;display:flex;flex-wrap:wrap;gap:4px'>{legend_items}</div>"
# Conversation bubbles
bubbles = ""
for seg in segments:
fg, bg = color_map[seg.speaker]
align = "flex-end" if speaker_list.index(seg.speaker) % 2 == 1 else "flex-start"
bubbles += f"""
<div style='display:flex;justify-content:{align};margin:6px 12px'>
<div style='max-width:75%'>
<div style='font-size:11px;color:#6b7280;margin-bottom:3px;
{"text-align:right" if align=="flex-end" else ""}'>
<span style='font-weight:600;color:{fg}'>Speaker {seg.speaker}</span>
&nbsp;·&nbsp;{seg.start_fmt}{seg.end_fmt}
</div>
<div style='background:{bg};border:1px solid {fg}30;color:#111827;
padding:10px 14px;border-radius:{"4px 16px 16px 16px" if align=="flex-start" else "16px 4px 16px 16px"};
font-size:14px;line-height:1.5'>
{seg.text}
</div>
</div>
</div>"""
return f"""
<div style='border:1px solid #e5e7eb;border-radius:12px;overflow:hidden;
font-family:system-ui,sans-serif;background:white'>
{legend}
<div style='height:480px;overflow-y:auto;padding:8px 0;background:#f9fafb'>
{bubbles}
</div>
</div>"""
def make_table_html(segments):
if not segments:
return ""
speaker_list = list(dict.fromkeys(s.speaker for s in segments))
color_map = {spk: SPEAKER_COLORS[i % len(SPEAKER_COLORS)] for i, spk in enumerate(speaker_list)}
rows = "".join(
f"<tr style='border-bottom:1px solid #f3f4f6'>"
f"<td style='padding:8px 12px'>"
f" <span style='background:{color_map[s.speaker][1]};color:{color_map[s.speaker][0]};"
f" padding:3px 10px;border-radius:99px;font-weight:700;font-size:13px'> {s.speaker}</span></td>"
f"<td style='padding:8px 12px;color:#6b7280;font-size:13px;white-space:nowrap'>{s.start_fmt}</td>"
f"<td style='padding:8px 12px;color:#6b7280;font-size:13px;white-space:nowrap'>{s.end_fmt}</td>"
f"<td style='padding:8px 12px;font-size:14px;color:#111827'>{s.text}</td>"
f"</tr>"
for s in segments
)
return f"""
<div style='border:1px solid #e5e7eb;border-radius:12px;overflow:hidden;font-family:system-ui,sans-serif'>
<table style='width:100%;border-collapse:collapse'>
<thead>
<tr style='background:#1e3a5f;color:white'>
<th style='padding:10px 12px;text-align:left;font-size:13px'>Speaker</th>
<th style='padding:10px 12px;text-align:left;font-size:13px'>Start</th>
<th style='padding:10px 12px;text-align:left;font-size:13px'>End</th>
<th style='padding:10px 12px;text-align:left;font-size:13px'>Transcript</th>
</tr>
</thead>
<tbody>{rows}</tbody>
</table>
</div>"""
def process_audio(audio_path, num_speakers):
global _last_segments
if audio_path is None:
return "⚠️ Please upload an audio file first.", "", ""
try:
pipeline = get_pipeline()
n = int(num_speakers) if num_speakers and int(num_speakers) > 0 else 0
segments = pipeline.process(audio_path, num_speakers=n)
except Exception as e:
return f"❌ Error: {e}\n{traceback.format_exc()}", "", ""
if not segments:
return "⚠️ No speech detected.", "", ""
_last_segments = segments
unique = len(set(s.speaker for s in segments))
status = f"✅ Done — {len(segments)} segments · {unique} speaker(s) detected"
return status, make_conversation_html(segments), make_table_html(segments)
def export_json():
if not _last_segments: return None
out = str(UPLOAD_DIR / "annotation.json")
to_json(_last_segments, out)
return out
def export_csv():
if not _last_segments: return None
out = str(UPLOAD_DIR / "annotation.csv")
to_csv(_last_segments, out)
return out
css = """
.gradio-container { max-width: 1100px !important; margin: auto !important; }
footer { display: none !important; }
"""
with gr.Blocks(title="Speech Annotation Pipeline", css=css) as demo:
gr.Markdown("""# 🎙️ Speech Annotation Pipeline
*Upload audio · Detect speakers · Export transcript*""")
with gr.Row():
with gr.Column(scale=1):
audio_input = gr.Audio(label="Upload Audio (.wav / .mp3 / .flac)", type="filepath")
num_speakers = gr.Slider(minimum=0, maximum=10, step=1, value=0,
label="Number of speakers (0 = auto-detect)")
run_btn = gr.Button("▶ Run Annotation", variant="primary", size="lg")
status_box = gr.Textbox(label="Status", value="Ready.", interactive=False)
gr.Markdown("### 📥 Export")
with gr.Row():
json_btn = gr.Button("⬇ JSON", size="sm")
csv_btn = gr.Button("⬇ CSV", size="sm")
json_file = gr.File(label="JSON Download", visible=True)
csv_file = gr.File(label="CSV Download", visible=True)
with gr.Column(scale=2):
gr.Markdown("### 💬 Conversation View")
conversation_html = gr.HTML(value="<div style='height:480px;border:1px solid #e5e7eb;"
"border-radius:12px;display:flex;align-items:center;"
"justify-content:center;color:#9ca3af;font-family:system-ui'>"
"Transcript will appear here after processing…</div>")
gr.Markdown("### 📋 Segment Table")
table_html = gr.HTML(value="")
run_btn.click(fn=process_audio,
inputs=[audio_input, num_speakers],
outputs=[status_box, conversation_html, table_html])
json_btn.click(fn=export_json, inputs=[], outputs=[json_file])
csv_btn.click(fn=export_csv, inputs=[], outputs=[csv_file])
if __name__ == "__main__":
demo.launch(server_name="0.0.0.0", server_port=7860, show_error=True)