diff --git a/README.md b/README.md index 80964e3..84b50d5 100644 --- a/README.md +++ b/README.md @@ -69,6 +69,9 @@ Prefer a clone and Poetry for development or source changes — see ## What you can do - **Create a BeWER report from text** — paste or upload reference/generated transcripts and Tympany runs bewer for you. *Generate report* produces the report (JSON) to view/download; *Generate & analyze* also runs the full classification. +- **Import Corti diarized transcripts** — paste a Corti transcript JSON (WebSocket streams or REST `/transcripts` format) and Tympany parses speaker labels, flattens the text for evaluation, and tags every error row with its speaker. +- **Per-speaker evaluation** — split a diarized transcript by speaker and compute WER/CER for each speaker individually, with per-speaker metrics shown alongside the overall figures. +- **Diarization error analysis** — align reference and generated speaker turns by time overlap to detect speaker mismatches, merged/split turns, and missing/extra turns. A **diarization accuracy** metric (matched turns / total turns) is shown alongside WER/CER. - **Analyze a BeWER report** — upload an existing BeWER report (JSON) and classify every difference. - **Medical Term Recall (MTR)** — supply a medical-terms list and bewer additionally reports how many of those terms the generated text captured. Save lists for reuse, or **extract** them automatically from a BeWER report. - **Review & iterate** — reclassify, exclude, flag, edit descriptions (all autosaved), then **Re-run** to measure the impact of excluded errors. Export results or flagged rows. @@ -85,6 +88,7 @@ The top nav has three sections: **Home** (landing page with these entry points p | **Replacement candidate** | Low | Abbreviation or shorthand fixable with a replacement rule | | **Context-dependent** | Medium | Spelling variant or compound boundary — meaning likely preserved, but needs review | | **Misrecognition** | High | Large edit distance errors that may have clinical significance or alter intended meaning | +| **Diarization error** | Medium/High | Speaker assignment problem — wrong speaker, merged/split turns, or missing/extra turns | --- @@ -193,9 +197,13 @@ Run flags for `./start.sh`: `HOST`, `PORT`, `WEB_CONCURRENCY`, and `TYMPANY_DEV= On the **Create** page: -- **Input** — paste reference and generated text (one example per line, paired by position) or upload a CSV with `ref`/`gen` columns. +- **Input** — three modes: + - **Paste text** — paste reference and generated text (one example per line, paired by position). + - **Upload CSV** — upload a CSV with `ref`/`gen` columns. + - **Import Corti transcript** — paste or upload a Corti transcript JSON. Both WebSocket streams format (`transcript`/`speakerId`/`participant.channel`/seconds) and REST `/transcripts` format (`text`/`channel`/milliseconds) are auto-detected. Speaker labels are parsed from `speakerId` and `channel`; text is flattened for evaluation and each edit row is tagged with its speaker. Use this to evaluate diarized transcripts directly without manually flattening them. - **Medical terms** *(optional)* — paste, upload a `.txt` (one term per line), or pick a saved list to compute Medical Term Recall. - **Normalization** — toggle text normalization before scoring. +- **Evaluate per speaker** *(visible with Corti transcript input)* — split the transcript by speaker and compute WER/CER for each speaker individually. Per-speaker metrics appear on the results page and in the BeWER report view. - **Generate report** runs bewer and keeps you on the page to view/download the report (JSON) and iterate; each generation is saved under **Generated BeWER reports**. - **Generate & analyze** runs bewer *and* the full Tympany classification, landing you on the results page. @@ -219,7 +227,9 @@ Supplying a medical-terms file makes bewer compute **Medical Term Recall (MTR)** On the results page each error is one editable row: reclassify it, edit its description, **Exclude** it (drops it from counts and from a re-run), or **Flag** it. Edits autosave to your history. -- **Export results** — every row in its current edited state. +When the transcript is diarized, each row also shows its **speaker** label, and a **Speaker** filter lets you view errors for a single speaker. The classification dropdown includes **Diarization error** as an option, and the breakdown table and filters include a diarization column. + +- **Export results** — every row in its current edited state (includes a `speaker` column when diarized). - **Export flags** — only flagged rows, with `ref_context`/`gen_context` columns holding the full reference/generated text of each flagged error's example. - For reports created in Tympany: **Download BeWER report** (the input report JSON) and **Download CSV** (the ref/gen source). The input and re-run reports can also be viewed inline from the links above the metrics. @@ -233,6 +243,44 @@ The results page shows **WER** and **CER** from the report alongside an **Update --- +## Diarized transcript support + +Tympany can import Corti transcripts that include speaker diarization metadata and evaluate them with speaker-aware analysis. + +### Importing Corti transcripts + +On the **Create** page, select **Import Corti transcript** to paste or upload a Corti transcript JSON. Tympany auto-detects the format: + +- **WebSocket streams** — segments arrive as `{ transcript, speakerId, participant: { channel }, time: { start, end } }` with times in seconds. Segments are sorted by start time. +- **REST `/transcripts`** — segments arrive as `{ text, speakerId, channel, start, end }` with times in milliseconds. +- **Minimal format** — a bare JSON array of `{ transcript, speakerId, channel }` objects (no `time`, `participant`, or `id` fields). Useful for reference transcripts that only need speaker labels, not timing. + +`speakerId` values of `0`–`3` indicate diarized speakers; `-1` means diarization is off (the segment is labeled by channel only). Channel is audio routing and is independent from speaker assignment. + +The parser (`tympany/diarize.py`) splits the transcript into **turns** — maximal runs of consecutive same-speaker segments. Each turn becomes a separate bewer example, so the report and analysis views show each speaker turn as its own example (e.g. "Example 1 — Speaker 0", "Example 2 — Speaker 1"). Turns are paired by position between reference and generated; if one side has more turns, the missing side is empty (full deletion or insertion). Every token, diff group, and error row on the results page is tagged with its speaker. + +### Per-speaker evaluation + +When **Evaluate per speaker** is checked (available with Corti transcript input), Tympany splits the transcript by speaker and runs bewer separately for each speaker's text. The results page shows a per-speaker metrics table (WER, CER, updated WER/CER, ref word count) alongside the overall figures. Re-run also recomputes per-speaker metrics. The BeWER report view shows alignment examples split by speaker. + +### Diarization error analysis + +When both the reference and generated sides are diarized, Tympany aligns speaker turns by time overlap and detects five types of diarization errors: + +| Error type | Description | Risk | +|---|---|---| +| **Speaker mismatch** | A generated segment covers the same time range as a reference segment but attributes the speech to a different speaker | Medium | +| **Merged turns** | Two or more reference segments (different speakers) are covered by a single generated segment | Medium | +| **Split turns** | One reference segment is covered by two or more generated segments attributed to different speakers | Medium | +| **Missing turn** | A reference segment has no overlapping generated segment | High | +| **Extra turn** | A generated segment has no overlapping reference segment | High | + +These appear as a **Diarization error** classification on the results page (with its own color, filter option, and breakdown column), alongside the existing misrecognition/formatting/context-dependent classifications. + +A **diarization accuracy** metric (matched turns / total turns) is computed and displayed on both the results page and the BeWER report view when diarization data is present. A turn is "matched" when its best-overlapping counterpart (coverage ≥ 50%) has the same `speakerId`. + +--- + ## How the analysis works ### Step 1 — Evaluation & alignment @@ -257,6 +305,8 @@ Each diff group passes through an ordered rule chain in `tympany/categorize.py`; | `rule_spelling_close` | Levenshtein distance ≤ 2 or similarity ≥ 0.8 | | `rule_other` | Fallback — `misrecognition` | +When diarization data is present, diarization errors (detected by `tympany/diarize_errors.py`) are classified separately as **diarization_error** — see [Diarized transcript support](#diarized-transcript-support). + Each edit also gets token ratio, character distance, and similarity score appended to its detail (shown as a tooltip on the Risk badge). ### Step 3 — LLM second pass (optional) @@ -290,12 +340,15 @@ tympany/ data.py — Medical word lists and abbreviations (EN, FR, DE, de-CH, DA) corti.py — Corti Agentic Framework client (classifier + term-extractor agents) ner.py — Optional local medical-entity NER backend (ALPHA, off by default) - bewer_eval.py — bewer evaluation (run_bewer) + re-run reconstruction + diarize.py — Parse Corti diarized transcript JSON (streams + REST) into speaker-tagged segments + diarize_errors.py — Align speaker turns by time overlap; detect 5 diarization error types + accuracy metric + bewer_eval.py — bewer evaluation (run_bewer, run_bewer_diarized) + re-run reconstruction web/ app.py — FastAPI application: auth, home, create, analyze, history, terms serve.py — web-server console entry point (tympany-serve) history.py — Per-user analysis + generated-report storage terms.py — Saved medical-term lists + report_render.py — BeWER report rendering (speaker-split examples, per-speaker metrics, diarization accuracy) templates/ — Jinja2 templates (base, login, home, create, analyze, results, reports_new, terms_edit, bewer_report) static/ — Static assets pyproject.toml — Python project metadata in Poetry format; poetry.lock pins versions @@ -333,6 +386,7 @@ corti_agent.json cached Corti agent ids (classifier, term_extra - The rule chain in `tympany/categorize.py` is ordered — first match wins. New rules go above `rule_other`. - `tympany/data.py` holds word lists per language. Add abbreviations or a new language's number words here before reaching for the LLM. - `parser.from_bewer` maps bewer alignment ops to diff groups; `tympany/bewer_eval.py` wraps the bewer library (pinned alpha — watch for API changes). +- `tympany/diarize.py` parses Corti transcript JSON (streams and REST formats) into `SpeakerSegment` objects; `tympany/diarize_errors.py` aligns speaker turns and detects diarization errors. bewer itself has no diarization support — all speaker logic is wrapped around flat-text bewer calls. ### Tests @@ -341,7 +395,7 @@ poetry install # includes the dev group (pytest) poetry run pytest ``` -The suite (`tests/`) covers the saved-term store, input validation, routing/redirects, and term extraction (the Corti agent call is mocked, so no credentials or network are needed). +The suite (`tests/`) covers the saved-term store, input validation, routing/redirects, term extraction, Corti transcript parsing, diarization error detection, per-speaker evaluation, and report rendering (the Corti agent call is mocked, so no credentials or network are needed). --- diff --git a/tests/fixtures/clinical_gen.json b/tests/fixtures/clinical_gen.json new file mode 100644 index 0000000..f160d5d --- /dev/null +++ b/tests/fixtures/clinical_gen.json @@ -0,0 +1,45 @@ +{ + "type": "transcript", + "data": [ + { + "id": "gen-0001", + "transcript": "Good morning, what brings you in today?", + "final": true, + "speakerId": 0, + "participant": { "channel": 0 }, + "time": { "start": 0.50, "end": 3.20 } + }, + { + "id": "gen-0002", + "transcript": "I have had a ever and a cough for 3 days.", + "final": true, + "speakerId": 1, + "participant": { "channel": 0 }, + "time": { "start": 3.50, "end": 7.80 } + }, + { + "id": "gen-0003", + "transcript": "I will prescribe you some antibiotic. Take them twice a day.", + "final": true, + "speakerId": 0, + "participant": { "channel": 0 }, + "time": { "start": 8.10, "end": 12.40 } + }, + { + "id": "gen-0004", + "transcript": "Thank you, doctor. When should I come back?", + "final": true, + "speakerId": 1, + "participant": { "channel": 0 }, + "time": { "start": 12.70, "end": 15.90 } + }, + { + "id": "gen-0005", + "transcript": "And drink plenty of fluids.", + "final": true, + "speakerId": 0, + "participant": { "channel": 0 }, + "time": { "start": 16.20, "end": 18.50 } + } + ] +} diff --git a/tests/fixtures/clinical_ref.json b/tests/fixtures/clinical_ref.json new file mode 100644 index 0000000..cd73a9f --- /dev/null +++ b/tests/fixtures/clinical_ref.json @@ -0,0 +1,7 @@ +[ + { "transcript": "Good morning, what brings you in today?", "speakerId": 0, "channel": 0 }, + { "transcript": "I have had a fever and a cough for three days.", "speakerId": 1, "channel": 0 }, + { "transcript": "I will prescribe you some antibiotics. Take them twice a day.", "speakerId": 0, "channel": 0 }, + { "transcript": "Thank you, doctor. When should I come back?", "speakerId": 1, "channel": 0 }, + { "transcript": "In two weeks if the symptoms persist.", "speakerId": 0, "channel": 0 } +] diff --git a/tests/fixtures/diarized_errors_gen.json b/tests/fixtures/diarized_errors_gen.json new file mode 100644 index 0000000..08cfa00 --- /dev/null +++ b/tests/fixtures/diarized_errors_gen.json @@ -0,0 +1,53 @@ +{ + "type": "transcript", + "data": [ + { + "id": "gen-0001", + "transcript": "Good morning, what brings you in today?", + "final": true, + "speakerId": 0, + "participant": { "channel": 0 }, + "time": { "start": 0.50, "end": 3.20 } + }, + { + "id": "gen-0002", + "transcript": "I have had a ever and a cough for 3 days.", + "final": true, + "speakerId": 0, + "participant": { "channel": 0 }, + "time": { "start": 3.50, "end": 7.80 } + }, + { + "id": "gen-0003a", + "transcript": "I will prescribe you some antibiotic.", + "final": true, + "speakerId": 0, + "participant": { "channel": 0 }, + "time": { "start": 8.10, "end": 11.00 } + }, + { + "id": "gen-0003b", + "transcript": "Take them twice a day.", + "final": true, + "speakerId": 1, + "participant": { "channel": 0 }, + "time": { "start": 11.00, "end": 12.40 } + }, + { + "id": "gen-0004", + "transcript": "Thank you, doctor. When should I come back?", + "final": true, + "speakerId": 1, + "participant": { "channel": 0 }, + "time": { "start": 12.70, "end": 15.90 } + }, + { + "id": "gen-0006", + "transcript": "And drink plenty of fluids.", + "final": true, + "speakerId": 0, + "participant": { "channel": 0 }, + "time": { "start": 19.00, "end": 21.00 } + } + ] +} diff --git a/tests/fixtures/diarized_errors_ref.json b/tests/fixtures/diarized_errors_ref.json new file mode 100644 index 0000000..42d41b3 --- /dev/null +++ b/tests/fixtures/diarized_errors_ref.json @@ -0,0 +1,45 @@ +{ + "type": "transcript", + "data": [ + { + "id": "ref-0001", + "transcript": "Good morning, what brings you in today?", + "final": true, + "speakerId": 0, + "participant": { "channel": 0 }, + "time": { "start": 0.50, "end": 3.20 } + }, + { + "id": "ref-0002", + "transcript": "I have had a fever and a cough for three days.", + "final": true, + "speakerId": 1, + "participant": { "channel": 0 }, + "time": { "start": 3.50, "end": 7.80 } + }, + { + "id": "ref-0003", + "transcript": "I will prescribe you some antibiotics. Take them twice a day.", + "final": true, + "speakerId": 0, + "participant": { "channel": 0 }, + "time": { "start": 8.10, "end": 12.40 } + }, + { + "id": "ref-0004", + "transcript": "Thank you, doctor. When should I come back?", + "final": true, + "speakerId": 1, + "participant": { "channel": 0 }, + "time": { "start": 12.70, "end": 15.90 } + }, + { + "id": "ref-0005", + "transcript": "In two weeks if the symptoms persist.", + "final": true, + "speakerId": 0, + "participant": { "channel": 0 }, + "time": { "start": 16.20, "end": 18.50 } + } + ] +} diff --git a/tests/fixtures/diarized_gen.json b/tests/fixtures/diarized_gen.json new file mode 100644 index 0000000..a5d3460 --- /dev/null +++ b/tests/fixtures/diarized_gen.json @@ -0,0 +1,37 @@ +{ + "type": "transcript", + "data": [ + { + "id": "b2c3d4e5-0000-0000-0000-000000000001", + "transcript": "Hello, what brings you in today?", + "final": true, + "speakerId": 0, + "participant": { "channel": 0 }, + "time": { "start": 0.40, "end": 3.10 } + }, + { + "id": "b2c3d4e5-0000-0000-0000-000000000002", + "transcript": "I've had a ever and a cough for 3 days.", + "final": true, + "speakerId": 1, + "participant": { "channel": 0 }, + "time": { "start": 3.40, "end": 7.20 } + }, + { + "id": "b2c3d4e5-0000-0000-0000-000000000003", + "transcript": "I'll prescribe you some antibiotic.", + "final": true, + "speakerId": 0, + "participant": { "channel": 0 }, + "time": { "start": 7.50, "end": 10.80 } + }, + { + "id": "b2c3d4e5-0000-0000-0000-000000000004", + "transcript": "Thank you, doctor.", + "final": true, + "speakerId": 1, + "participant": { "channel": 0 }, + "time": { "start": 11.00, "end": 12.50 } + } + ] +} diff --git a/tests/fixtures/diarized_ref.json b/tests/fixtures/diarized_ref.json new file mode 100644 index 0000000..b71182a --- /dev/null +++ b/tests/fixtures/diarized_ref.json @@ -0,0 +1,37 @@ +{ + "type": "transcript", + "data": [ + { + "id": "a1b2c3d4-0000-0000-0000-000000000001", + "transcript": "Hello, what brings you in today?", + "final": true, + "speakerId": 0, + "participant": { "channel": 0 }, + "time": { "start": 0.40, "end": 3.10 } + }, + { + "id": "a1b2c3d4-0000-0000-0000-000000000002", + "transcript": "I've had a fever and a cough for three days.", + "final": true, + "speakerId": 1, + "participant": { "channel": 0 }, + "time": { "start": 3.40, "end": 7.20 } + }, + { + "id": "a1b2c3d4-0000-0000-0000-000000000003", + "transcript": "I'll prescribe you some antibiotics.", + "final": true, + "speakerId": 0, + "participant": { "channel": 0 }, + "time": { "start": 7.50, "end": 10.80 } + }, + { + "id": "a1b2c3d4-0000-0000-0000-000000000004", + "transcript": "Thank you, doctor.", + "final": true, + "speakerId": 1, + "participant": { "channel": 0 }, + "time": { "start": 11.00, "end": 12.50 } + } + ] +} diff --git a/tests/test_bewer.py b/tests/test_bewer.py index 74748c1..c4cb9a5 100644 --- a/tests/test_bewer.py +++ b/tests/test_bewer.py @@ -35,6 +35,51 @@ def test_from_bewer_builds_samples_and_diffs(): assert samples[1].diffs == () # identical example → no diffs +def test_from_bewer_tags_tokens_with_speakers(): + ref_ws = [["Speaker 0", "Speaker 0", "Speaker 0", "Speaker 0", "Speaker 1"]] + gen_ws = [["Speaker 0", "Speaker 0", "Speaker 0", "Speaker 1"]] + samples = from_bewer(ENVELOPE, ref_word_speakers=ref_ws, gen_word_speakers=gen_ws) + assert len(samples) == 2 + ex = samples[0] + assert ex.ref_tokens[0].speaker == "Speaker 0" + assert ex.ref_tokens[4].speaker == "Speaker 1" + assert ex.pred_tokens[0].speaker == "Speaker 0" + assert ex.pred_tokens[3].speaker == "Speaker 1" + assert "Speaker 1" in ex.speakers + assert "Speaker 0" in ex.speakers + assert ex.diffs + assert ex.diffs[0].speaker + + +def test_from_bewer_reads_speakers_from_envelope(): + import copy + env = copy.deepcopy(ENVELOPE) + env["examples"][0]["ref_speakers"] = ["Speaker 0"] * 5 + env["examples"][0]["hyp_speakers"] = ["Speaker 0"] * 4 + samples = from_bewer(env) + assert samples[0].ref_tokens[0].speaker == "Speaker 0" + assert "Speaker 0" in samples[0].speakers + + +def test_from_bewer_no_speakers_backward_compat(): + samples = from_bewer(ENVELOPE) + assert all(t.speaker == "" for s in samples for t in s.ref_tokens) + assert all(t.speaker == "" for s in samples for t in s.pred_tokens) + assert all(s.speakers == () for s in samples) + + +def test_samples_to_payload_includes_speaker(): + from tympany.parser import samples_to_payload + ref_ws = [["Speaker 0", "Speaker 0", "Speaker 0", "Speaker 0", "Speaker 1"]] + gen_ws = [["Speaker 0", "Speaker 0", "Speaker 0", "Speaker 1"]] + samples = from_bewer(ENVELOPE, ref_word_speakers=ref_ws, gen_word_speakers=gen_ws) + payload = samples_to_payload(samples) + assert "speaker" in payload[0]["ref_tokens"][0] + assert "speakers" in payload[0] + assert payload[0]["ref_tokens"][0]["speaker"] == "Speaker 0" + assert "Speaker 1" in payload[0]["speakers"] + + def test_from_bewer_diffs_feed_categorizer(): samples = from_bewer(ENVELOPE) edits = [e for d in samples[0].diffs for e in categorize_group(d)] @@ -44,12 +89,55 @@ def test_from_bewer_diffs_feed_categorizer(): } +def test_categorize_group_preserves_speaker(): + from tympany.parser import DiffGroup + group = DiffGroup(("hello",), ("helloo",), speaker="Speaker 1") + edits = categorize_group(group) + assert edits + assert edits[0].speaker == "Speaker 1" + + def test_metrics_from_bewer(): m = metrics_from_bewer(ENVELOPE) assert _PCT.fullmatch(m["wer"]) and _PCT.fullmatch(m["cer"]) and _PCT.fullmatch(m["mtr"]) assert m["normalization"] is True +def test_metrics_from_bewer_includes_per_speaker(): + env = { + "metrics": {"wer": "10.00%", "cer": "5.00%", "mtr": None}, + "settings": {"normalization": True, "diarized": True}, + "per_speaker": { + "Speaker 0": { + "metrics": {"wer": "0.00%", "cer": "0.00%", "mtr": None}, + "ref_words": 6, + "gen_words": 6, + }, + }, + } + m = metrics_from_bewer(env) + assert "per_speaker" in m + assert m["per_speaker"]["Speaker 0"]["wer"] == "0.00%" + assert m["per_speaker"]["Speaker 0"]["ref_words"] == 6 + + +def test_metrics_from_bewer_no_per_speaker(): + m = metrics_from_bewer(ENVELOPE) + assert "per_speaker" not in m + + +def test_metrics_from_bewer_includes_diarization_accuracy(): + env = { + "metrics": {"wer": "10.00%", "cer": "5.00%", "mtr": None}, + "settings": {"normalization": True, "diarized": True}, + "diarization_accuracy": {"accuracy": "75.00%", "matched": 3, "total": 4}, + } + m = metrics_from_bewer(env) + assert m["diarization_accuracy"] == "75.00%" + assert m["diarization_matched"] == 3 + assert m["diarization_total"] == 4 + + def test_reference_corpus_from_bewer_samples(): samples = from_bewer(ENVELOPE) assert reference_corpus(samples) == "the patient has hypertension today\nblood pressure was elevated" @@ -81,3 +169,139 @@ def test_run_bewer_live_honors_normalization_flag(): assert [op["type"] for op in raw["examples"][0]["ops"]] == [ "SUBSTITUTE", "SUBSTITUTE", ] + + +def test_run_bewer_diarized_live(): + bewer = pytest.importorskip("bewer") + from tympany.bewer_eval import run_bewer_diarized + from tympany.diarize import SpeakerSegment + + ref_segs = [ + SpeakerSegment(0, 0, "Hello what brings you in today", 0.4, 3.1), + SpeakerSegment(1, 0, "I have had a fever", 3.4, 7.2), + ] + gen_segs = [ + SpeakerSegment(0, 0, "Hello what brings you in today", 0.4, 3.1), + SpeakerSegment(1, 0, "I have had a ever", 3.4, 7.2), + ] + env = run_bewer_diarized(ref_segs, gen_segs) + + assert "per_speaker" in env + assert "Speaker 0" in env["per_speaker"] + assert "Speaker 1" in env["per_speaker"] + assert env["settings"]["diarized"] is True + + spk0 = env["per_speaker"]["Speaker 0"] + spk1 = env["per_speaker"]["Speaker 1"] + assert spk0["metrics"]["wer"] == "0.00%" + assert spk1["metrics"]["wer"] != "0.00%" + assert spk0["ref_words"] == 6 + assert spk1["ref_words"] == 5 + + +def test_build_rows_per_speaker(): + from tympany.bewer_eval import build_rows_per_speaker + + samples = [{ + "example": 1, + "ref_tokens": [ + {"cls": "ok", "text": "hello", "speaker": "Speaker 0"}, + {"cls": "sub", "text": "fever", "speaker": "Speaker 1"}, + ], + "pred_tokens": [ + {"cls": "ok", "text": "hello", "speaker": "Speaker 0"}, + {"cls": "sub", "text": "ever", "speaker": "Speaker 1"}, + ], + }] + edits = [{ + "example": 1, "speaker": "Speaker 1", + "ref": "fever", "gen": "ever", + "excluded": True, + }] + speaker_rows = build_rows_per_speaker(samples, edits) + assert "Speaker 0" in speaker_rows + assert "Speaker 1" in speaker_rows + # Speaker 0 had no errors → (ref, gen) identical + spk0 = speaker_rows["Speaker 0"][0] + assert spk0[0] == "hello" + assert spk0[1] == "hello" + # Speaker 1's error was excluded → corrected to ref + spk1 = speaker_rows["Speaker 1"][0] + assert spk1[0] == "fever" + assert spk1[1] == "fever" + + +def test_run_bewer_diarized_includes_accuracy(): + bewer = pytest.importorskip("bewer") + from tympany.bewer_eval import run_bewer_diarized + from tympany.diarize import SpeakerSegment + + ref_segs = [ + SpeakerSegment(0, 0, "hello doctor", 0.0, 3.0), + SpeakerSegment(1, 0, "I am fine", 3.0, 6.0), + ] + gen_segs = [ + SpeakerSegment(0, 0, "hello doctor", 0.0, 3.0), + SpeakerSegment(0, 0, "I am fine", 3.0, 6.0), + ] + env = run_bewer_diarized(ref_segs, gen_segs) + assert env["diarization_accuracy"] is not None + assert env["diarization_accuracy"]["accuracy"] == "50.00%" + assert env["diarization_accuracy"]["matched"] == 1 + assert env["diarization_accuracy"]["total"] == 2 + + +def test_run_bewer_diarized_produces_per_turn_examples(): + bewer = pytest.importorskip("bewer") + from tympany.bewer_eval import run_bewer_diarized + from tympany.diarize import SpeakerSegment + + ref_segs = [ + SpeakerSegment(0, 0, "Hello what brings you in today", 0.4, 3.1), + SpeakerSegment(1, 0, "I have had a fever", 3.4, 7.2), + SpeakerSegment(0, 0, "I will prescribe antibiotics", 7.5, 10.8), + ] + gen_segs = [ + SpeakerSegment(0, 0, "Hello what brings you in today", 0.4, 3.1), + SpeakerSegment(1, 0, "I have had a ever", 3.4, 7.2), + SpeakerSegment(0, 0, "I will prescribe antibiotic", 7.5, 10.8), + ] + env = run_bewer_diarized(ref_segs, gen_segs) + + # Three turns (alternating speakers) → three examples + assert len(env["examples"]) == 3 + assert env["examples"][0]["example"] == 1 + assert env["examples"][1]["example"] == 2 + assert env["examples"][2]["example"] == 3 + + # Each example's speaker labels should be uniform within the example + ref_spk_0 = env["examples"][0]["ref_speakers"] + assert all(s == "Speaker 0" for s in ref_spk_0) + ref_spk_1 = env["examples"][1]["ref_speakers"] + assert all(s == "Speaker 1" for s in ref_spk_1) + + # Per-speaker metrics still present + assert "Speaker 0" in env["per_speaker"] + assert "Speaker 1" in env["per_speaker"] + + +def test_run_bewer_diarized_unequal_turn_counts(): + bewer = pytest.importorskip("bewer") + from tympany.bewer_eval import run_bewer_diarized + from tympany.diarize import SpeakerSegment + + ref_segs = [ + SpeakerSegment(0, 0, "Hello doctor", 0, 2), + SpeakerSegment(1, 0, "Hi", 2, 3), + SpeakerSegment(0, 0, "Goodbye", 3, 5), + ] + gen_segs = [ + SpeakerSegment(0, 0, "Hello doctor", 0, 2), + SpeakerSegment(1, 0, "Hi", 2, 3), + ] + env = run_bewer_diarized(ref_segs, gen_segs) + + # Ref has 3 turns, gen has 2 → 3 examples (last is a full deletion) + assert len(env["examples"]) == 3 + assert env["examples"][2]["ref"] == "Goodbye" + assert env["examples"][2]["hyp"] == "" diff --git a/tests/test_diarize.py b/tests/test_diarize.py new file mode 100644 index 0000000..94f1e6b --- /dev/null +++ b/tests/test_diarize.py @@ -0,0 +1,315 @@ +"""Tests for Corti diarized transcript parsing.""" + +import json + +import pytest + +from tympany.diarize import ( + SpeakerSegment, + Turn, + build_word_speakers, + distinct_speakers, + flatten_segments, + group_segments_by_turn, + is_diarized, + parse_corti_transcript, + parse_corti_transcript_json, + split_into_turns, +) + + +STREAMS_MSG = { + "type": "transcript", + "data": [ + { + "id": "uuid-1", + "transcript": "Hello, what brings you in today?", + "final": True, + "speakerId": 0, + "participant": {"channel": 0}, + "time": {"start": 6.50, "end": 8.90}, + }, + { + "id": "uuid-2", + "transcript": "I've had a fever and a cough.", + "final": True, + "speakerId": 1, + "participant": {"channel": 0}, + "time": {"start": 3.40, "end": 6.20}, + }, + ], +} + +REST_MSG = { + "id": "f47ac10b-58cc-4372-a567-0e02b2c3d479", + "metadata": { + "participantsRoles": [ + {"channel": 0, "role": "doctor"}, + {"channel": 1, "role": "patient"}, + ] + }, + "transcripts": [ + {"channel": 0, "participant": 0, "speakerId": 0, + "text": "Hello, what brings you in today?", + "start": 400, "end": 3100}, + {"channel": 1, "participant": 1, "speakerId": 1, + "text": "I've had a fever and a cough.", + "start": 3400, "end": 6200}, + ], + "usageInfo": {"creditsConsumed": 0.42}, + "recordingId": "abc12300-0000-0000-0000-000000000001", + "status": "completed", +} + +NO_DIARIZ_MSG = { + "type": "transcript", + "data": [ + { + "id": "uuid-0", + "transcript": "Patient presents with fever and cough.", + "final": True, + "speakerId": -1, + "participant": {"channel": 0}, + "time": {"start": 1.71, "end": 11.296}, + } + ], +} + + +def test_parse_streams_sorts_by_start_time(): + segments = parse_corti_transcript(STREAMS_MSG) + assert len(segments) == 2 + assert segments[0].start < segments[1].start + assert segments[0].text == "I've had a fever and a cough." + assert segments[1].text == "Hello, what brings you in today?" + + +def test_parse_streams_fields(): + segments = parse_corti_transcript(STREAMS_MSG) + seg = segments[0] + assert seg.speaker_id == 1 + assert seg.channel == 0 + assert seg.start == 3.40 + assert seg.end == 6.20 + + +def test_parse_rest_converts_ms_to_seconds(): + segments = parse_corti_transcript(REST_MSG) + assert len(segments) == 2 + assert segments[0].start == 0.400 + assert segments[0].end == 3.100 + assert segments[0].text == "Hello, what brings you in today?" + assert segments[0].speaker_id == 0 + assert segments[0].channel == 0 + + +def test_parse_rest_second_segment(): + segments = parse_corti_transcript(REST_MSG) + seg = segments[1] + assert seg.speaker_id == 1 + assert seg.channel == 1 + assert seg.start == 3.400 + assert seg.end == 6.200 + + +def test_parse_json_string(): + raw = json.dumps(STREAMS_MSG) + segments = parse_corti_transcript_json(raw) + assert len(segments) == 2 + + +def test_parse_json_array_treated_as_streams(): + raw = json.dumps(STREAMS_MSG["data"]) + segments = parse_corti_transcript_json(raw) + assert len(segments) == 2 + + +def test_parse_invalid_json(): + with pytest.raises(ValueError, match="Invalid JSON"): + parse_corti_transcript_json("{not json}") + + +def test_parse_unrecognized_format(): + with pytest.raises(ValueError, match="Unrecognized Corti transcript format"): + parse_corti_transcript({"foo": "bar"}) + + +def test_speaker_label_diarized(): + seg = SpeakerSegment(speaker_id=0, channel=0, text="hi", start=0, end=1) + assert seg.label == "Speaker 0" + + +def test_speaker_label_no_diarization(): + seg = SpeakerSegment(speaker_id=-1, channel=1, text="hi", start=0, end=1) + assert seg.label == "Channel 1" + + +def test_flatten_segments(): + segments = parse_corti_transcript(STREAMS_MSG) + text = flatten_segments(segments) + assert "fever" in text + assert "Hello" in text + assert text == "I've had a fever and a cough. Hello, what brings you in today?" + + +def test_flatten_segments_skips_empty(): + segments = [ + SpeakerSegment(0, 0, "", 0, 1), + SpeakerSegment(0, 0, "hello", 1, 2), + ] + assert flatten_segments(segments) == "hello" + + +def test_build_word_speakers(): + segments = parse_corti_transcript(STREAMS_MSG) + labels = build_word_speakers(segments) + assert len(labels) == 13 + assert labels[0] == "Speaker 1" + assert labels[7] == "Speaker 0" + + +def test_is_diarized_true(): + segments = parse_corti_transcript(STREAMS_MSG) + assert is_diarized(segments) is True + + +def test_is_diarized_false(): + segments = parse_corti_transcript(NO_DIARIZ_MSG) + assert is_diarized(segments) is False + + +def test_distinct_speakers(): + segments = parse_corti_transcript(STREAMS_MSG) + speakers = distinct_speakers(segments) + assert speakers == ["Speaker 1", "Speaker 0"] + + +def test_distinct_speakers_no_diarization(): + segments = parse_corti_transcript(NO_DIARIZ_MSG) + speakers = distinct_speakers(segments) + assert speakers == ["Channel 0"] + + +# --------------------------------------------------------------------------- +# split_into_turns +# --------------------------------------------------------------------------- + +def test_split_into_turns_alternating_speakers(): + segments = [ + SpeakerSegment(0, 0, "Hello doctor", 0, 2), + SpeakerSegment(1, 0, "Hi there", 2, 4), + SpeakerSegment(0, 0, "How are you", 4, 6), + ] + turns = split_into_turns(segments) + assert len(turns) == 3 + assert turns[0] == Turn("Speaker 0", "Hello doctor") + assert turns[1] == Turn("Speaker 1", "Hi there") + assert turns[2] == Turn("Speaker 0", "How are you") + + +def test_split_into_turns_merges_consecutive_same_speaker(): + segments = [ + SpeakerSegment(0, 0, "Hello", 0, 1), + SpeakerSegment(0, 0, "doctor", 1, 2), + SpeakerSegment(1, 0, "Hi", 2, 3), + ] + turns = split_into_turns(segments) + assert len(turns) == 2 + assert turns[0] == Turn("Speaker 0", "Hello doctor") + assert turns[1] == Turn("Speaker 1", "Hi") + + +def test_split_into_turns_non_diarized_single_turn(): + segments = [ + SpeakerSegment(-1, 0, "Hello", 0, 1), + SpeakerSegment(-1, 0, "world", 1, 2), + ] + turns = split_into_turns(segments) + assert len(turns) == 1 + assert turns[0] == Turn("Channel 0", "Hello world") + + +def test_split_into_turns_skips_empty_segments(): + segments = [ + SpeakerSegment(0, 0, "Hello", 0, 1), + SpeakerSegment(0, 0, "", 1, 2), + SpeakerSegment(1, 0, "Hi", 2, 3), + ] + turns = split_into_turns(segments) + assert len(turns) == 2 + assert turns[0] == Turn("Speaker 0", "Hello") + assert turns[1] == Turn("Speaker 1", "Hi") + + +def test_split_into_turns_empty_input(): + assert split_into_turns([]) == [] + + +def test_parse_streams_top_level_channel(): + """Minimal ref format: channel at top level, no participant wrapper.""" + data = {"type": "transcript", "data": [ + {"transcript": "Hello", "speakerId": 0, "channel": 1}, + ]} + segments = parse_corti_transcript(data) + assert segments[0].channel == 1 + + +# --------------------------------------------------------------------------- +# group_segments_by_turn +# --------------------------------------------------------------------------- + +def test_group_segments_by_turn_alternating(): + segments = [ + SpeakerSegment(0, 0, "Hello", 0, 1), + SpeakerSegment(1, 0, "Hi", 1, 2), + SpeakerSegment(0, 0, "Bye", 2, 3), + ] + groups = group_segments_by_turn(segments) + assert len(groups) == 3 + assert [s.text for s in groups[0]] == ["Hello"] + assert [s.text for s in groups[1]] == ["Hi"] + assert [s.text for s in groups[2]] == ["Bye"] + + +def test_group_segments_by_turn_merges_consecutive(): + segments = [ + SpeakerSegment(0, 0, "Hello", 0, 1), + SpeakerSegment(0, 0, "doctor", 1, 2), + SpeakerSegment(1, 0, "Hi", 2, 3), + ] + groups = group_segments_by_turn(segments) + assert len(groups) == 2 + assert [s.text for s in groups[0]] == ["Hello", "doctor"] + assert [s.text for s in groups[1]] == ["Hi"] + + +def test_group_segments_by_turn_skips_empty(): + segments = [ + SpeakerSegment(0, 0, "Hello", 0, 1), + SpeakerSegment(0, 0, "", 1, 2), + SpeakerSegment(1, 0, "Hi", 2, 3), + ] + groups = group_segments_by_turn(segments) + assert len(groups) == 2 + assert [s.text for s in groups[0]] == ["Hello"] + assert [s.text for s in groups[1]] == ["Hi"] + + +def test_group_segments_by_turn_empty(): + assert group_segments_by_turn([]) == [] + + +def test_group_segments_matches_split_into_turns(): + """group_segments_by_turn and split_into_turns must produce the same grouping.""" + segments = [ + SpeakerSegment(0, 0, "Hello doctor", 0, 2), + SpeakerSegment(1, 0, "Hi there", 2, 4), + SpeakerSegment(1, 0, "how are you", 4, 6), + SpeakerSegment(0, 0, "Goodbye", 6, 8), + ] + turns = split_into_turns(segments) + groups = group_segments_by_turn(segments) + assert len(turns) == len(groups) + for turn, group in zip(turns, groups): + assert turn.speaker == group[0].label + assert turn.text == " ".join(s.text.strip() for s in group) diff --git a/tests/test_diarize_errors.py b/tests/test_diarize_errors.py new file mode 100644 index 0000000..b3dd8e8 --- /dev/null +++ b/tests/test_diarize_errors.py @@ -0,0 +1,161 @@ +"""Tests for diarization error detection (tympany.diarize_errors).""" + +from tympany.diarize import SpeakerSegment +from tympany.diarize_errors import align_segments + + +def _seg(speaker_id, channel, text, start, end): + return SpeakerSegment(speaker_id, channel, text, start, end) + + +def test_no_errors_when_segments_match(): + ref = [ + _seg(0, 0, "hello doctor", 0.4, 3.1), + _seg(1, 0, "I have a fever", 3.4, 7.2), + ] + gen = [ + _seg(0, 0, "hello doctor", 0.4, 3.1), + _seg(1, 0, "I have a fever", 3.4, 7.2), + ] + errors = align_segments(ref, gen) + assert errors == [] + + +def test_speaker_mismatch(): + ref = [ + _seg(0, 0, "hello doctor", 0.4, 3.1), + _seg(1, 0, "I have a fever", 3.4, 7.2), + ] + gen = [ + _seg(1, 0, "hello doctor", 0.4, 3.1), + _seg(0, 0, "I have a fever", 3.4, 7.2), + ] + errors = align_segments(ref, gen) + mismatches = [e for e in errors if e.category == "speaker_mismatch"] + assert len(mismatches) == 2 + assert "Speaker 0" in mismatches[0].detail or "Speaker 1" in mismatches[0].detail + + +def test_merged_turns(): + ref = [ + _seg(0, 0, "how are you", 0.0, 3.0), + _seg(1, 0, "I am fine", 3.0, 6.0), + ] + gen = [ + _seg(0, 0, "how are you I am fine", 0.0, 6.0), + ] + errors = align_segments(ref, gen) + merges = [e for e in errors if e.category == "speaker_merge"] + assert len(merges) == 1 + assert "Merged turns" in merges[0].detail + assert "how are you" in merges[0].ref_text + assert "I am fine" in merges[0].ref_text + + +def test_split_turn(): + ref = [ + _seg(0, 0, "how are you I am fine", 0.0, 6.0), + ] + gen = [ + _seg(0, 0, "how are you", 0.0, 3.0), + _seg(1, 0, "I am fine", 3.0, 6.0), + ] + errors = align_segments(ref, gen) + splits = [e for e in errors if e.category == "speaker_split"] + assert len(splits) == 1 + assert "Split turn" in splits[0].detail + assert "how are you I am fine" in splits[0].ref_text + + +def test_missing_turn(): + ref = [ + _seg(0, 0, "hello doctor", 0.4, 3.1), + _seg(1, 0, "I have a fever", 3.4, 7.2), + _seg(0, 0, "get well soon", 8.0, 10.0), + ] + gen = [ + _seg(0, 0, "hello doctor", 0.4, 3.1), + _seg(1, 0, "I have a fever", 3.4, 7.2), + ] + errors = align_segments(ref, gen) + missing = [e for e in errors if e.category == "missing_turn"] + assert len(missing) == 1 + assert "get well soon" in missing[0].ref_text + assert missing[0].gen_text == "" + + +def test_extra_turn(): + ref = [ + _seg(0, 0, "hello doctor", 0.4, 3.1), + ] + gen = [ + _seg(0, 0, "hello doctor", 0.4, 3.1), + _seg(1, 0, "extra words here", 4.0, 7.0), + ] + errors = align_segments(ref, gen) + extra = [e for e in errors if e.category == "extra_turn"] + assert len(extra) == 1 + assert "extra words here" in extra[0].gen_text + assert extra[0].ref_text == "" + + +def test_empty_segments_returns_no_errors(): + assert align_segments([], []) == [] + + +def test_diarization_error_carries_speaker(): + ref = [_seg(0, 0, "hi", 0, 1)] + gen = [_seg(1, 0, "hi", 0, 1)] + errors = align_segments(ref, gen) + assert errors + assert errors[0].speaker != "" + + +# --------------------------------------------------------------------------- +# Diarization accuracy +# --------------------------------------------------------------------------- + +def test_diarization_accuracy_perfect(): + from tympany.diarize_errors import diarization_accuracy + ref = [ + _seg(0, 0, "hello doctor", 0.0, 3.0), + _seg(1, 0, "I am fine", 3.0, 6.0), + ] + gen = [ + _seg(0, 0, "hello doctor", 0.0, 3.0), + _seg(1, 0, "I am fine", 3.0, 6.0), + ] + acc, matched, total = diarization_accuracy(ref, gen) + assert acc == 1.0 + assert matched == 2 + assert total == 2 + + +def test_diarization_accuracy_with_mismatch(): + from tympany.diarize_errors import diarization_accuracy + ref = [ + _seg(0, 0, "hello doctor", 0.0, 3.0), + _seg(1, 0, "I am fine", 3.0, 6.0), + ] + gen = [ + _seg(0, 0, "hello doctor", 0.0, 3.0), + _seg(0, 0, "I am fine", 3.0, 6.0), + ] + acc, matched, total = diarization_accuracy(ref, gen) + assert acc == 0.5 + assert matched == 1 + assert total == 2 + + +def test_diarization_accuracy_no_diarization(): + from tympany.diarize_errors import diarization_accuracy + ref = [_seg(-1, 0, "hi", 0, 1)] + gen = [_seg(-1, 0, "hi", 0, 1)] + acc, matched, total = diarization_accuracy(ref, gen) + assert acc is None + + +def test_diarization_accuracy_empty(): + from tympany.diarize_errors import diarization_accuracy + acc, matched, total = diarization_accuracy([], []) + assert acc is None diff --git a/tests/test_report_render.py b/tests/test_report_render.py index 31f554a..a8fc88e 100644 --- a/tests/test_report_render.py +++ b/tests/test_report_render.py @@ -120,3 +120,117 @@ def test_adjacent_terms_get_separate_boxes(): ops = [_op("MATCH", "nausea", "nausea"), _op("MATCH", "dizziness", "dizziness")] view = report_render.canal_view(_envelope(ops, key_terms=["nausea", "dizziness"])) assert str(view["examples"][0]["lines"][0]["ref"]).count("keyword-box") == 2 + + +# --------------------------------------------------------------------------- +# Per-speaker splitting (Phase 2) +# --------------------------------------------------------------------------- + +def _envelope_with_speakers(ops, ref_spk, hyp_spk, *, key_terms=None): + env = _envelope(ops, key_terms=key_terms) + env["examples"][0]["ref_speakers"] = ref_spk + env["examples"][0]["hyp_speakers"] = hyp_spk + return env + + +def test_split_ops_by_speaker_two_speakers(): + ops = [ + _op("MATCH", "hello", "hello"), + _op("MATCH", "doctor", "doctor"), + _op("MATCH", "I", "I"), + _op("MATCH", "see", "see"), + ] + ref_spk = ["Speaker 0", "Speaker 0", "Speaker 1", "Speaker 1"] + hyp_spk = ["Speaker 0", "Speaker 0", "Speaker 1", "Speaker 1"] + ex = {"ops": ops, "ref_speakers": ref_spk, "hyp_speakers": hyp_spk} + groups = report_render._split_ops_by_speaker(ex) + assert len(groups) == 2 + assert groups[0][0] == "Speaker 0" + assert len(groups[0][1]) == 2 + assert groups[1][0] == "Speaker 1" + assert len(groups[1][1]) == 2 + + +def test_canal_view_splits_examples_by_speaker(): + ops = [ + _op("MATCH", "hello", "hello"), + _op("MATCH", "doctor", "doctor"), + _op("MATCH", "I", "I"), + _op("MATCH", "see", "see"), + ] + ref_spk = ["Speaker 0", "Speaker 0", "Speaker 1", "Speaker 1"] + hyp_spk = ["Speaker 0", "Speaker 0", "Speaker 1", "Speaker 1"] + view = report_render.canal_view(_envelope_with_speakers(ops, ref_spk, hyp_spk)) + assert len(view["examples"]) == 2 + assert view["examples"][0]["speaker"] == "Speaker 0" + assert view["examples"][1]["speaker"] == "Speaker 1" + assert len(view["examples"][0]["lines"]) > 0 + assert len(view["examples"][1]["lines"]) > 0 + + +def test_canal_view_no_speakers_stays_single_example(): + ops = [_op("MATCH", "hello", "hello"), _op("MATCH", "doctor", "doctor")] + view = report_render.canal_view(_envelope(ops)) + assert len(view["examples"]) == 1 + assert view["examples"][0]["speaker"] == "" + + +def test_canal_view_speaker_split_preserves_word_counts(): + """Splitting by speaker doesn't change the summary corpus counts.""" + ops = [ + _op("MATCH", "hello", "hello"), + _op("MATCH", "doctor", "doctor"), + _op("MATCH", "I", "I"), + _op("MATCH", "see", "see"), + ] + ref_spk = ["Speaker 0", "Speaker 0", "Speaker 1", "Speaker 1"] + hyp_spk = ["Speaker 0", "Speaker 0", "Speaker 1", "Speaker 1"] + view = report_render.canal_view(_envelope_with_speakers(ops, ref_spk, hyp_spk)) + assert view["summary"]["ref_words"] == "4" + assert view["summary"]["gen_words"] == "4" + + +def test_canal_view_per_speaker_metrics(): + env = _envelope([_op("MATCH", "hi", "hi")]) + env["per_speaker"] = { + "Speaker 0": { + "metrics": {"wer": "0.00%", "cer": "0.00%", "mtr": None}, + "ref_words": 2, + "gen_words": 2, + }, + "Speaker 1": { + "metrics": {"wer": "50.00%", "cer": "25.00%", "mtr": None}, + "ref_words": 4, + "gen_words": 4, + }, + } + view = report_render.canal_view(env) + assert view["show_per_speaker"] is True + rows = view["per_speaker_rows"] + assert len(rows) == 2 + assert rows[0]["label"] == "Speaker 0" + assert rows[0]["wer"] == "0.00%" + assert rows[0]["ref_words"] == "2" + assert rows[1]["label"] == "Speaker 1" + assert rows[1]["wer"] == "50.00%" + assert rows[1]["cer"] == "25.00%" + + +def test_canal_view_no_per_speaker_when_absent(): + view = report_render.canal_view(_envelope([_op("MATCH", "hi", "hi")])) + assert view["show_per_speaker"] is False + assert view["per_speaker_rows"] == [] + + +def test_canal_view_diarization_accuracy(): + env = _envelope([_op("MATCH", "hi", "hi")]) + env["diarization_accuracy"] = {"accuracy": "80.00%", "matched": 4, "total": 5} + view = report_render.canal_view(env) + assert view["show_diarization_accuracy"] is True + assert view["diarization_accuracy"]["accuracy"] == "80.00%" + assert view["diarization_accuracy"]["matched"] == 4 + + +def test_canal_view_no_diarization_accuracy_when_absent(): + view = report_render.canal_view(_envelope([_op("MATCH", "hi", "hi")])) + assert view["show_diarization_accuracy"] is False diff --git a/tympany/bewer_eval.py b/tympany/bewer_eval.py index 0724ceb..adb2983 100644 --- a/tympany/bewer_eval.py +++ b/tympany/bewer_eval.py @@ -24,7 +24,10 @@ from __future__ import annotations -from typing import Optional, Sequence +from typing import Optional, Sequence, TYPE_CHECKING + +if TYPE_CHECKING: + from tympany.diarize import SpeakerSegment class BewerError(RuntimeError): @@ -46,11 +49,16 @@ def run_bewer( *, normalization: bool = True, medical_terms: Optional[Sequence[str]] = None, + ref_word_speakers: Optional[list[list[str]]] = None, + gen_word_speakers: Optional[list[list[str]]] = None, ) -> dict: """Evaluate (ref, gen) ``rows`` with bewer and return the JSON envelope. ``medical_terms`` enables the key-term-found (MTR) metric. ``normalization`` is recorded in settings; bewer applies its default standardisation pipeline. + ``ref_word_speakers`` and ``gen_word_speakers`` (one list per example, one + speaker label per word) are embedded in the envelope so downstream + consumers (report renderer, parser) can tag tokens with their speaker. """ try: from bewer import Dataset @@ -85,6 +93,10 @@ def run_bewer( ex = ds.examples[i] ops = [op.to_dict() for op in align.get_example_metric(ex).alignment] entry = {"example": i + 1, "ref": ref, "hyp": gen, "ops": ops} + if ref_word_speakers and i < len(ref_word_speakers): + entry["ref_speakers"] = ref_word_speakers[i] + if gen_word_speakers and i < len(gen_word_speakers): + entry["hyp_speakers"] = gen_word_speakers[i] # Record where bewer located each key term (as token-index [start, stop) # slices, aligned 1:1 with the ops' tokens) so the report can box exactly # the terms MTR counts — no second, divergent matcher. Best-effort: bewer @@ -206,3 +218,127 @@ def build_rows(samples: list[dict], edits: list[dict]) -> list[tuple[str, str]]: ) rows.append((ref, gen)) return rows + + +# --------------------------------------------------------------------------- +# Per-speaker evaluation (Phase 2) +# --------------------------------------------------------------------------- + +def run_bewer_diarized( + ref_segments: "list[SpeakerSegment]", + gen_segments: "list[SpeakerSegment]", + *, + normalization: bool = True, + medical_terms: Optional[Sequence[str]] = None, +) -> dict: + """Evaluate diarized transcripts per turn and per speaker. + + Splits each side into turns (consecutive same-speaker segments), pairs + them by position, and evaluates each pair as a separate bewer example. + Returns the same JSON envelope as ``run_bewer`` (with speaker-tagged + tokens), augmented with a ``per_speaker`` dict mapping each speaker label + to its own metrics + word counts and a ``diarization_accuracy`` block. + """ + from .diarize import ( + Turn, + distinct_speakers, + flatten_segments, + is_diarized, + split_into_turns, + ) + + ref_turns = split_into_turns(ref_segments) + gen_turns = split_into_turns(gen_segments) + n = max(len(ref_turns), len(gen_turns)) + + rows: list[tuple[str, str]] = [] + ref_ws: list[list[str]] = [] + gen_ws: list[list[str]] = [] + for i in range(n): + rt = ref_turns[i] if i < len(ref_turns) else Turn("", "") + gt = gen_turns[i] if i < len(gen_turns) else Turn("", "") + rows.append((rt.text, gt.text)) + ref_ws.append([rt.speaker] * len(rt.text.split()) if rt.text else []) + gen_ws.append([gt.speaker] * len(gt.text.split()) if gt.text else []) + + envelope = run_bewer( + rows, + normalization=normalization, + medical_terms=medical_terms, + ref_word_speakers=ref_ws, + gen_word_speakers=gen_ws, + ) + + per_speaker: dict[str, dict] = {} + if is_diarized(ref_segments) or is_diarized(gen_segments): + for spk in distinct_speakers(ref_segments + gen_segments): + spk_ref = flatten_segments([s for s in ref_segments if s.label == spk]) + spk_gen = flatten_segments([s for s in gen_segments if s.label == spk]) + if not spk_ref and not spk_gen: + continue + try: + spk_env = run_bewer( + [(spk_ref, spk_gen)], + normalization=normalization, + medical_terms=medical_terms, + ) + except BewerError: + continue + per_speaker[spk] = { + "metrics": spk_env["metrics"], + "ref_words": len(spk_ref.split()), + "gen_words": len(spk_gen.split()), + } + + envelope["per_speaker"] = per_speaker + envelope["settings"]["diarized"] = True + + # Diarization accuracy (Phase 3): matched turns / total turns. + from .diarize_errors import diarization_accuracy + acc, matched, total = diarization_accuracy(ref_segments, gen_segments) + envelope["diarization_accuracy"] = { + "accuracy": _pct(acc), + "matched": matched, + "total": total, + } if acc is not None else None + + return envelope + + +def build_rows_per_speaker( + samples: list[dict], edits: list[dict] +) -> dict[str, list[tuple[str, str]]]: + """Reconstruct per-speaker (ref, corrected_gen) rows for a diarized re-run. + + Groups tokens within each sample by their ``speaker`` field and + reconstructs corrected text independently per speaker, applying only + that speaker's excluded edits. Returns a dict mapping speaker label → + list of (ref, gen) rows (one per sample where that speaker appears). + """ + edits_by_example: dict[str, list[dict]] = {} + for edit in edits: + edits_by_example.setdefault(str(edit.get("example", "")), []).append(edit) + + speaker_rows: dict[str, list[tuple[str, str]]] = {} + + for sample in samples: + example = str(sample.get("example", "")) + ref_tokens = sample.get("ref_tokens", []) + pred_tokens = sample.get("pred_tokens", []) + sample_edits = edits_by_example.get(example, []) + + speakers: list[str] = [] + for t in ref_tokens + pred_tokens: + spk = t.get("speaker", "") + if spk and spk not in speakers: + speakers.append(spk) + + for spk in speakers: + spk_ref = [t for t in ref_tokens if t.get("speaker") == spk] + spk_pred = [t for t in pred_tokens if t.get("speaker") == spk] + spk_edits = [e for e in sample_edits if e.get("speaker") == spk] + ref, gen = reconstruct_example(spk_ref, spk_pred, spk_edits) + if ref.strip() or gen.strip(): + speaker_rows.setdefault(spk, []).append((ref, gen)) + + return speaker_rows diff --git a/tympany/categorize.py b/tympany/categorize.py index c99232f..a4dcf71 100644 --- a/tympany/categorize.py +++ b/tympany/categorize.py @@ -3,10 +3,11 @@ Each diff group is classified into a category, then mapped to: classification — one of: formatting_error | replacement_candidate | - context_dependent | misrecognition - risk_level — low | medium | high - replacement_candidate — True if the error looks like a fixable STT - replacement/command rule + context_dependent | misrecognition | + diarization_error + risk_level — low | medium | high + replacement_candidate — True if the error looks like a fixable STT + replacement/command rule Category → classification → risk mapping ----------------------------------------- @@ -21,6 +22,9 @@ misrecognition (medium/high risk — true errors): misrecognition, medication_or_device, pure_insertion, pure_deletion + +diarization_error (medium/high risk — speaker assignment errors): + speaker_mismatch, speaker_merge, speaker_split, missing_turn, extra_turn """ from __future__ import annotations @@ -66,6 +70,11 @@ "medication_or_device": ("misrecognition", "high"), "pure_insertion": ("misrecognition", "medium"), "pure_deletion": ("misrecognition", "medium"), + "speaker_mismatch": ("diarization_error", "medium"), + "speaker_merge": ("diarization_error", "medium"), + "speaker_split": ("diarization_error", "medium"), + "missing_turn": ("diarization_error", "high"), + "extra_turn": ("diarization_error", "high"), } @@ -79,6 +88,7 @@ class Edit: pred: tuple[str, ...] category: str detail: str = "" + speaker: str = "" @property def op(self) -> str: @@ -442,14 +452,15 @@ def classify( ref: tuple[str, ...], pred: tuple[str, ...], medical_terms: frozenset[str] = frozenset(), + speaker: str = "", ) -> Edit: for rule in _effective_rules(medical_terms): result = rule(ref, pred) if result is not None: category, detail = result enriched = (detail + _edit_stats(ref, pred)).strip() - return Edit(ref=ref, pred=pred, category=category, detail=enriched) - return Edit(ref=ref, pred=pred, category="misrecognition", detail=_edit_stats(ref, pred).strip()) + return Edit(ref=ref, pred=pred, category=category, detail=enriched, speaker=speaker) + return Edit(ref=ref, pred=pred, category="misrecognition", detail=_edit_stats(ref, pred).strip(), speaker=speaker) def _split_group(group: DiffGroup) -> list[tuple[tuple[str, ...], tuple[str, ...]]]: @@ -466,4 +477,5 @@ def _split_group(group: DiffGroup) -> list[tuple[tuple[str, ...], tuple[str, ... def categorize_group( group: DiffGroup, medical_terms: frozenset[str] = frozenset() ) -> list[Edit]: - return [classify(r, p, medical_terms) for r, p in _split_group(group)] + speaker = group.speaker + return [classify(r, p, medical_terms, speaker=speaker) for r, p in _split_group(group)] diff --git a/tympany/classify.py b/tympany/classify.py index 81718c2..5f408dd 100644 --- a/tympany/classify.py +++ b/tympany/classify.py @@ -26,10 +26,47 @@ def _sample_medical_terms(sample) -> frozenset[str]: return detect_entities(text) +def _diarization_error_rows( + sample, ref_segments, gen_segments, +) -> list[dict]: + """Detect diarization errors for a sample and return edit-row dicts. + + Only runs when both ref and gen segments are available (diarized input). + The returned rows carry the same shape as word-level edit rows so they + flow through ``normalize_edits`` and the results table unchanged. + """ + if not ref_segments or not gen_segments: + return [] + + from .diarize_errors import align_segments + from .categorize import _CATEGORY_META + + errors = align_segments(ref_segments, gen_segments) + rows: list[dict] = [] + for err in errors: + classification, risk = _CATEGORY_META.get(err.category, ("diarization_error", "medium")) + rows.append({ + "file": sample.file_stem, + "example": sample.example_num, + "ref": err.ref_text, + "gen": err.gen_text, + "op": "sub", + "category": err.category, + "classification": classification, + "risk_level": risk, + "replacement_candidate": False, + "detail": err.detail, + "speaker": err.speaker, + }) + return rows + + def classify_samples( samples, llm_provider: Optional[str] = None, llm_outcome: Optional[MutableMapping] = None, + ref_segments: Optional[list] = None, + gen_segments: Optional[list] = None, ) -> list[dict]: """Categorize every diff group across ``samples`` into edit-row dicts. @@ -38,10 +75,14 @@ def classify_samples( When the LLM pass runs and ``llm_outcome`` is provided, it is populated with the pass result (see ``tympany.llm.classify_with_llm``) so callers can tell a real LLM classification apart from a silent rule-based fallback. + + When ``ref_segments`` and ``gen_segments`` are provided (diarized input), + diarization errors (speaker mismatch, merged/split turns, missing/extra + turns) are detected and appended as additional edit rows. """ provider = detect_provider(llm_provider) if llm_provider else None rows: list[dict] = [] - for sample in samples: + for idx, sample in enumerate(samples): medical_terms = _sample_medical_terms(sample) edits = [ edit @@ -62,5 +103,10 @@ def classify_samples( "risk_level": edit.risk_level, "replacement_candidate": edit.is_replacement_candidate, "detail": edit.detail, + "speaker": edit.speaker, }) + + if ref_segments and gen_segments and idx < len(ref_segments) and idx < len(gen_segments): + rows.extend(_diarization_error_rows(sample, ref_segments[idx], gen_segments[idx])) + return rows diff --git a/tympany/diarize.py b/tympany/diarize.py new file mode 100644 index 0000000..a030a52 --- /dev/null +++ b/tympany/diarize.py @@ -0,0 +1,170 @@ +"""Parse Corti diarized transcript JSON into speaker-tagged segments. + +Corti exposes two transcript shapes: + +- **Streams** (WebSocket): ``{ type: "transcript", data: [{ id, transcript, + final, speakerId, participant: { channel }, time: { start, end } }] }`` + Segments arrive out of order; must be sorted by ``time.start``. + +- **REST** (``/transcripts``): ``{ id, metadata: { participantsRoles }, + transcripts: [{ channel, participant, speakerId, text, start, end }], ... }`` + Already ordered; times are in milliseconds. + +``parse_corti_transcript`` auto-detects the format. ``flatten_segments`` joins +segment texts into a single string for bewer evaluation. ``build_word_speakers`` +produces a per-word speaker label list so ``parser.from_bewer`` can tag each +token with its speaker. +""" + +from __future__ import annotations + +import json +from dataclasses import dataclass + + +@dataclass(frozen=True) +class SpeakerSegment: + speaker_id: int + channel: int + text: str + start: float + end: float + + @property + def label(self) -> str: + if self.speaker_id < 0: + return f"Channel {self.channel}" + return f"Speaker {self.speaker_id}" + + +def _detect_format(data: dict) -> str: + if data.get("type") == "transcript": + return "streams" + if "transcripts" in data: + return "rest" + raise ValueError( + "Unrecognized Corti transcript format — expected a streams message " + '(type: "transcript") or a REST response (with a "transcripts" key).' + ) + + +def _parse_streams(data: dict) -> list[SpeakerSegment]: + segments: list[SpeakerSegment] = [] + for seg in data.get("data", []): + speaker_id = int(seg.get("speakerId", -1)) + participant = seg.get("participant") or {} + channel = int(participant.get("channel", seg.get("channel", 0))) + text = seg.get("transcript") or "" + time = seg.get("time") or {} + start = float(time.get("start", 0)) + end = float(time.get("end", 0)) + segments.append(SpeakerSegment(speaker_id, channel, text, start, end)) + segments.sort(key=lambda s: (s.start, s.end)) + return segments + + +def _parse_rest(data: dict) -> list[SpeakerSegment]: + segments: list[SpeakerSegment] = [] + for seg in data.get("transcripts") or []: + speaker_id = int(seg.get("speakerId", -1)) + channel = int(seg.get("channel", 0)) + text = seg.get("text") or "" + start = float(seg.get("start", 0)) / 1000.0 + end = float(seg.get("end", 0)) / 1000.0 + segments.append(SpeakerSegment(speaker_id, channel, text, start, end)) + return segments + + +def parse_corti_transcript(data: dict) -> list[SpeakerSegment]: + fmt = _detect_format(data) + if fmt == "streams": + return _parse_streams(data) + return _parse_rest(data) + + +def parse_corti_transcript_json(raw: str) -> list[SpeakerSegment]: + try: + data = json.loads(raw) + except ValueError as exc: + raise ValueError(f"Invalid JSON: {exc}") from exc + if isinstance(data, list): + data = {"type": "transcript", "data": data} + if not isinstance(data, dict): + raise ValueError("Expected a JSON object or array of transcript segments.") + return parse_corti_transcript(data) + + +@dataclass(frozen=True) +class Turn: + """A consecutive run of same-speaker segments treated as one evaluation unit.""" + speaker: str + text: str + + +def split_into_turns(segments: list[SpeakerSegment]) -> list[Turn]: + """Group consecutive same-speaker segments into turns. + + A *turn* is a maximal run of segments whose :attr:`SpeakerSegment.label` + is the same. Non-diarized transcripts (all ``speakerId == -1``) produce a + single turn, preserving the old flat-evaluation behaviour. + """ + turns: list[Turn] = [] + cur_speaker = "" + cur_parts: list[str] = [] + for seg in segments: + if not seg.text.strip(): + continue + if seg.label != cur_speaker: + if cur_parts: + turns.append(Turn(cur_speaker, " ".join(cur_parts))) + cur_speaker = seg.label + cur_parts = [seg.text.strip()] + else: + cur_parts.append(seg.text.strip()) + if cur_parts: + turns.append(Turn(cur_speaker, " ".join(cur_parts))) + return turns + + +def group_segments_by_turn(segments: list[SpeakerSegment]) -> list[list[SpeakerSegment]]: + """Group consecutive same-speaker segments into lists, matching ``split_into_turns``. + + Returns one list of raw segments per turn (skipping empty-text segments). + The grouping matches :func:`split_into_turns` so callers can correlate + per-turn text with the underlying segments for diarization error analysis. + """ + groups: list[list[SpeakerSegment]] = [] + cur_speaker = "" + for seg in segments: + if not seg.text.strip(): + continue + if seg.label != cur_speaker: + groups.append([seg]) + cur_speaker = seg.label + else: + groups[-1].append(seg) + return groups + + +def flatten_segments(segments: list[SpeakerSegment]) -> str: + return " ".join(s.text for s in segments if s.text.strip()) + + +def build_word_speakers(segments: list[SpeakerSegment]) -> list[str]: + labels: list[str] = [] + for seg in segments: + for _ in seg.text.split(): + labels.append(seg.label) + return labels + + +def is_diarized(segments: list[SpeakerSegment]) -> bool: + return any(s.speaker_id >= 0 for s in segments) + + +def distinct_speakers(segments: list[SpeakerSegment]) -> list[str]: + seen: list[str] = [] + for seg in segments: + if seg.label not in seen: + seen.append(seg.label) + return seen diff --git a/tympany/diarize_errors.py b/tympany/diarize_errors.py new file mode 100644 index 0000000..eaa3917 --- /dev/null +++ b/tympany/diarize_errors.py @@ -0,0 +1,355 @@ +"""Detect diarization errors by aligning reference and generated speaker turns. + +Given time-ordered speaker segments from both sides, we align them by time +overlap and check for: + +- **Speaker mismatch** — a gen segment covers the same time range as a ref + segment but attributes the speech to a different speaker. +- **Merged turns** — two or more ref segments (different speakers) are + covered by a single gen segment. +- **Split turns** — one ref segment is covered by two or more gen segments + attributed to different speakers. +- **Missing turn** — a ref segment has no overlapping gen segment. +- **Extra turn** — a gen segment has no overlapping ref segment. + +Each error is returned as a :class:`DiarizationError` with a category +matching the ``_CATEGORY_META`` table in ``categorize.py``. +""" + +from __future__ import annotations + +from dataclasses import dataclass + +from .diarize import SpeakerSegment + + +@dataclass(frozen=True) +class DiarizationError: + category: str + detail: str + speaker: str + ref_text: str + gen_text: str + start: float + end: float + + +def _overlap(a: SpeakerSegment, b: SpeakerSegment) -> float: + """Time overlap in seconds between two segments.""" + return max(0.0, min(a.end, b.end) - max(a.start, b.start)) + + +def _coverage(seg: SpeakerSegment, other: SpeakerSegment) -> float: + """Fraction of ``seg``'s duration covered by ``other``.""" + dur = seg.end - seg.start + if dur <= 0: + return 0.0 + return _overlap(seg, other) / dur + + +def align_segments( + ref_segments: list[SpeakerSegment], + gen_segments: list[SpeakerSegment], + *, + coverage_threshold: float = 0.5, +) -> list[DiarizationError]: + """Align ref and gen speaker turns by time overlap and detect diarization errors. + + Returns a list of :class:`DiarizationError` instances. When the segments + are well-aligned (same speakers, same turns, no missing/extra), the list + is empty. + """ + errors: list[DiarizationError] = [] + if not ref_segments or not gen_segments: + return errors + + # For each ref segment, find the best-matching gen segment by time overlap. + used_gen: set[int] = set() + ref_to_gen: dict[int, list[int]] = {} + + for ri, rseg in enumerate(ref_segments): + best_overlap = 0.0 + best_gi = -1 + for gi, gseg in enumerate(gen_segments): + ov = _overlap(rseg, gseg) + if ov > best_overlap: + best_overlap = ov + best_gi = gi + if best_gi >= 0 and _coverage(rseg, gen_segments[best_gi]) >= coverage_threshold: + ref_to_gen.setdefault(best_gi, []).append(ri) + + # For each gen segment, find the best-matching ref segment (reverse direction). + gen_to_ref: dict[int, list[int]] = {} + for gi, gseg in enumerate(gen_segments): + best_overlap = 0.0 + best_ri = -1 + for ri, rseg in enumerate(ref_segments): + ov = _overlap(gseg, rseg) + if ov > best_overlap: + best_overlap = ov + best_ri = ri + if best_ri >= 0 and _coverage(gseg, ref_segments[best_ri]) >= coverage_threshold: + gen_to_ref.setdefault(best_ri, []).append(gi) + + # Detect: merged turns (multiple ref → one gen), split turns (one ref → multiple gen), + # and speaker mismatches (1:1 match but different speakers). + for gi, ref_indices in ref_to_gen.items(): + gseg = gen_segments[gi] + if len(ref_indices) > 1: + ref_speakers = {ref_segments[ri].speaker_id for ri in ref_indices} + if len(ref_speakers) > 1: + ref_texts = [ref_segments[ri].text for ri in ref_indices] + errors.append(DiarizationError( + category="speaker_merge", + detail=f"Merged turns: {len(ref_indices)} ref segments → 1 gen segment", + speaker=gseg.label, + ref_text=" | ".join(ref_texts), + gen_text=gseg.text, + start=min(ref_segments[ri].start for ri in ref_indices), + end=max(ref_segments[ri].end for ri in ref_indices), + )) + used_gen.add(gi) + + for ri, gen_indices in gen_to_ref.items(): + rseg = ref_segments[ri] + if len(gen_indices) > 1: + gen_speakers = {gen_segments[gi].speaker_id for gi in gen_indices} + if len(gen_speakers) > 1: + gen_texts = [gen_segments[gi].text for gi in gen_indices] + errors.append(DiarizationError( + category="speaker_split", + detail=f"Split turn: 1 ref segment → {len(gen_indices)} gen segments", + speaker=rseg.label, + ref_text=rseg.text, + gen_text=" | ".join(gen_texts), + start=min(gen_segments[gi].start for gi in gen_indices), + end=max(gen_segments[gi].end for gi in gen_indices), + )) + for gi in gen_indices: + used_gen.add(gi) + + # Check 1:1 matches for speaker mismatch. + for gi, ref_indices in ref_to_gen.items(): + if len(ref_indices) == 1 and gi not in used_gen: + ri = ref_indices[0] + rseg = ref_segments[ri] + gseg = gen_segments[gi] + if rseg.speaker_id >= 0 and gseg.speaker_id >= 0 and rseg.speaker_id != gseg.speaker_id: + errors.append(DiarizationError( + category="speaker_mismatch", + detail=f"Speaker mismatch: ref {rseg.label} → gen {gseg.label}", + speaker=gseg.label, + ref_text=rseg.text, + gen_text=gseg.text, + start=rseg.start, + end=rseg.end, + )) + used_gen.add(gi) + + # Missing turns: ref segments not matched to any gen segment. + matched_refs: set[int] = set() + for ref_indices in ref_to_gen.values(): + matched_refs.update(ref_indices) + for ri, rseg in enumerate(ref_segments): + if ri not in matched_refs: + errors.append(DiarizationError( + category="missing_turn", + detail=f"Missing turn: no generated segment for ref {rseg.label}", + speaker=rseg.label, + ref_text=rseg.text, + gen_text="", + start=rseg.start, + end=rseg.end, + )) + + # Extra turns: gen segments not matched to any ref segment. + for gi, gseg in enumerate(gen_segments): + if gi not in used_gen: + matched_as_secondary = any(gi in gen_indices for gen_indices in gen_to_ref.values()) + if not matched_as_secondary: + errors.append(DiarizationError( + category="extra_turn", + detail=f"Extra turn: no reference segment for gen {gseg.label}", + speaker=gseg.label, + ref_text="", + gen_text=gseg.text, + start=gseg.start, + end=gseg.end, + )) + + return errors + + +def diarization_accuracy( + ref_segments: list[SpeakerSegment], + gen_segments: list[SpeakerSegment], + *, + coverage_threshold: float = 0.5, +) -> "tuple[Optional[float], int, int]": + """Compute diarization accuracy: matched turns / total turns. + + A turn is "matched" when its best-overlapping counterpart (coverage >= + threshold) has the same speakerId. Returns ``(accuracy, matched, total)`` + where ``accuracy`` is a 0–1 float (or ``None`` when neither side is + diarized or there are no turns to compare). + """ + if not ref_segments or not gen_segments: + return None, 0, 0 + if not any(s.speaker_id >= 0 for s in ref_segments + gen_segments): + return None, 0, 0 + + matched = 0 + total = 0 + used_gen: set[int] = set() + + for rseg in ref_segments: + best_gi = -1 + best_cov = 0.0 + for gi, gseg in enumerate(gen_segments): + cov = _coverage(rseg, gseg) + if cov > best_cov: + best_cov = cov + best_gi = gi + if best_gi >= 0 and best_cov >= coverage_threshold: + total += 1 + gseg = gen_segments[best_gi] + if rseg.speaker_id >= 0 and gseg.speaker_id >= 0 and rseg.speaker_id == gseg.speaker_id: + matched += 1 + used_gen.add(best_gi) + + for gi, gseg in enumerate(gen_segments): + if gi not in used_gen: + best_ri = -1 + best_cov = 0.0 + for ri, rseg in enumerate(ref_segments): + cov = _coverage(gseg, rseg) + if cov > best_cov: + best_cov = cov + best_ri = ri + if best_ri >= 0 and best_cov >= coverage_threshold: + total += 1 + rseg = ref_segments[best_ri] + if gseg.speaker_id >= 0 and rseg.speaker_id >= 0 and gseg.speaker_id == rseg.speaker_id: + matched += 1 + + if total == 0: + return None, 0, 0 + return matched / total, matched, total + + # For each ref segment, find the best-matching gen segment by time overlap. + used_gen: set[int] = set() + ref_to_gen: dict[int, list[int]] = {} + + for ri, rseg in enumerate(ref_segments): + best_overlap = 0.0 + best_gi = -1 + for gi, gseg in enumerate(gen_segments): + ov = _overlap(rseg, gseg) + if ov > best_overlap: + best_overlap = ov + best_gi = gi + if best_gi >= 0 and _coverage(rseg, gen_segments[best_gi]) >= coverage_threshold: + ref_to_gen.setdefault(best_gi, []).append(ri) + + # For each gen segment, find the best-matching ref segment (reverse direction). + gen_to_ref: dict[int, list[int]] = {} + for gi, gseg in enumerate(gen_segments): + best_overlap = 0.0 + best_ri = -1 + for ri, rseg in enumerate(ref_segments): + ov = _overlap(gseg, rseg) + if ov > best_overlap: + best_overlap = ov + best_ri = ri + if best_ri >= 0 and _coverage(gseg, ref_segments[best_ri]) >= coverage_threshold: + gen_to_ref.setdefault(best_ri, []).append(gi) + + # Detect: merged turns (multiple ref → one gen), split turns (one ref → multiple gen), + # and speaker mismatches (1:1 match but different speakers). + for gi, ref_indices in ref_to_gen.items(): + gseg = gen_segments[gi] + if len(ref_indices) > 1: + ref_speakers = {ref_segments[ri].speaker_id for ri in ref_indices} + if len(ref_speakers) > 1: + ref_texts = [ref_segments[ri].text for ri in ref_indices] + errors.append(DiarizationError( + category="speaker_merge", + detail=f"Merged turns: {len(ref_indices)} ref segments → 1 gen segment", + speaker=gseg.label, + ref_text=" | ".join(ref_texts), + gen_text=gseg.text, + start=min(ref_segments[ri].start for ri in ref_indices), + end=max(ref_segments[ri].end for ri in ref_indices), + )) + used_gen.add(gi) + + for ri, gen_indices in gen_to_ref.items(): + rseg = ref_segments[ri] + if len(gen_indices) > 1: + gen_speakers = {gen_segments[gi].speaker_id for gi in gen_indices} + if len(gen_speakers) > 1: + gen_texts = [gen_segments[gi].text for gi in gen_indices] + errors.append(DiarizationError( + category="speaker_split", + detail=f"Split turn: 1 ref segment → {len(gen_indices)} gen segments", + speaker=rseg.label, + ref_text=rseg.text, + gen_text=" | ".join(gen_texts), + start=min(gen_segments[gi].start for gi in gen_indices), + end=max(gen_segments[gi].end for gi in gen_indices), + )) + for gi in gen_indices: + used_gen.add(gi) + + # Check 1:1 matches for speaker mismatch. + for gi, ref_indices in ref_to_gen.items(): + if len(ref_indices) == 1 and gi not in used_gen: + ri = ref_indices[0] + rseg = ref_segments[ri] + gseg = gen_segments[gi] + if rseg.speaker_id >= 0 and gseg.speaker_id >= 0 and rseg.speaker_id != gseg.speaker_id: + errors.append(DiarizationError( + category="speaker_mismatch", + detail=f"Speaker mismatch: ref {rseg.label} → gen {gseg.label}", + speaker=gseg.label, + ref_text=rseg.text, + gen_text=gseg.text, + start=rseg.start, + end=rseg.end, + )) + used_gen.add(gi) + + # Missing turns: ref segments not matched to any gen segment. + matched_refs: set[int] = set() + for ref_indices in ref_to_gen.values(): + matched_refs.update(ref_indices) + for ri, rseg in enumerate(ref_segments): + if ri not in matched_refs: + errors.append(DiarizationError( + category="missing_turn", + detail=f"Missing turn: no generated segment for ref {rseg.label}", + speaker=rseg.label, + ref_text=rseg.text, + gen_text="", + start=rseg.start, + end=rseg.end, + )) + + # Extra turns: gen segments not matched to any ref segment. + for gi, gseg in enumerate(gen_segments): + if gi not in used_gen: + if gi not in {gi2 for ref_indices in ref_to_gen.values() for gi2 in [gi]}: + pass + matched_as_secondary = any(gi in gen_indices for gen_indices in gen_to_ref.values()) + if not matched_as_secondary: + errors.append(DiarizationError( + category="extra_turn", + detail=f"Extra turn: no reference segment for gen {gseg.label}", + speaker=gseg.label, + ref_text="", + gen_text=gseg.text, + start=gseg.start, + end=gseg.end, + )) + + return errors diff --git a/tympany/parser.py b/tympany/parser.py index 87d475a..58c4f11 100644 --- a/tympany/parser.py +++ b/tympany/parser.py @@ -8,6 +8,7 @@ from __future__ import annotations from dataclasses import dataclass +from typing import Optional # --------------------------------------------------------------------------- @@ -18,12 +19,14 @@ class Token: cls: str # "ok" | "sub" | "del" | "ins" text: str + speaker: str = "" @dataclass(frozen=True) class DiffGroup: ref: tuple[str, ...] pred: tuple[str, ...] + speaker: str = "" @property def op(self) -> str: @@ -41,6 +44,7 @@ class Sample: ref_tokens: tuple[Token, ...] pred_tokens: tuple[Token, ...] diffs: tuple[DiffGroup, ...] + speakers: tuple[str, ...] = () def samples_to_payload(samples: list[Sample]) -> list[dict]: @@ -53,75 +57,142 @@ def samples_to_payload(samples: list[Sample]) -> list[dict]: return [ { "example": s.example_num, - "ref_tokens": [{"cls": t.cls, "text": t.text} for t in s.ref_tokens], - "pred_tokens": [{"cls": t.cls, "text": t.text} for t in s.pred_tokens], + "ref_tokens": [{"cls": t.cls, "text": t.text, "speaker": t.speaker} for t in s.ref_tokens], + "pred_tokens": [{"cls": t.cls, "text": t.text, "speaker": t.speaker} for t in s.pred_tokens], + "speakers": list(s.speakers), } for s in samples ] -def from_bewer(envelope: dict, file_stem: str = "report") -> list[Sample]: +def from_bewer( + envelope: dict, + file_stem: str = "report", + ref_word_speakers: Optional[list[list[str]]] = None, + gen_word_speakers: Optional[list[list[str]]] = None, +) -> list[Sample]: """Map a bewer JSON envelope (see tympany.bewer_eval) into Samples. Each alignment op becomes ref/pred Tokens (MATCH→ok, SUBSTITUTE→sub, DELETE→del, INSERT→ins); contiguous non-match ops are grouped into DiffGroups, exactly the shape the categorizer consumes. This is the This is how Tympany turns a bewer evaluation into reviewable diffs. + + When ``ref_word_speakers`` and ``gen_word_speakers`` are provided (one + list per example, one speaker label per word), tokens are tagged with + their speaker and each DiffGroup carries the speaker of its first token. """ samples: list[Sample] = [] - for ex in envelope.get("examples", []): + for ex_idx, ex in enumerate(envelope.get("examples", [])): ref_tokens: list[Token] = [] pred_tokens: list[Token] = [] diffs: list[DiffGroup] = [] cur_ref: list[str] = [] cur_pred: list[str] = [] + cur_speaker: str = "" + + ref_ws = ref_word_speakers[ex_idx] if ref_word_speakers and ex_idx < len(ref_word_speakers) else ex.get("ref_speakers") + gen_ws = gen_word_speakers[ex_idx] if gen_word_speakers and ex_idx < len(gen_word_speakers) else ex.get("hyp_speakers") + ref_idx = 0 + gen_idx = 0 def _flush() -> None: + nonlocal cur_speaker if cur_ref or cur_pred: - diffs.append(DiffGroup(tuple(cur_ref), tuple(cur_pred))) + diffs.append(DiffGroup(tuple(cur_ref), tuple(cur_pred), speaker=cur_speaker)) cur_ref.clear() cur_pred.clear() + cur_speaker = "" for op in ex.get("ops", []): op_type = (op.get("type") or "").upper() ref, hyp = op.get("ref"), op.get("hyp") if op_type == "MATCH": _flush() - ref_tokens.append(Token("ok", ref)) - pred_tokens.append(Token("ok", hyp)) + ref_spk = ref_ws[ref_idx] if ref_ws and ref_idx < len(ref_ws) else "" + gen_spk = gen_ws[gen_idx] if gen_ws and gen_idx < len(gen_ws) else "" + ref_tokens.append(Token("ok", ref, speaker=ref_spk)) + pred_tokens.append(Token("ok", hyp, speaker=gen_spk)) + ref_idx += 1 + gen_idx += 1 elif op_type == "SUBSTITUTE": - ref_tokens.append(Token("sub", ref)) - pred_tokens.append(Token("sub", hyp)) + ref_spk = ref_ws[ref_idx] if ref_ws and ref_idx < len(ref_ws) else "" + gen_spk = gen_ws[gen_idx] if gen_ws and gen_idx < len(gen_ws) else "" + if not cur_speaker: + cur_speaker = ref_spk or gen_spk + ref_tokens.append(Token("sub", ref, speaker=ref_spk)) + pred_tokens.append(Token("sub", hyp, speaker=gen_spk)) cur_ref.append(ref) cur_pred.append(hyp) + ref_idx += 1 + gen_idx += 1 elif op_type == "DELETE": - ref_tokens.append(Token("del", ref)) + ref_spk = ref_ws[ref_idx] if ref_ws and ref_idx < len(ref_ws) else "" + if not cur_speaker: + cur_speaker = ref_spk + ref_tokens.append(Token("del", ref, speaker=ref_spk)) cur_ref.append(ref) + ref_idx += 1 elif op_type == "INSERT": - pred_tokens.append(Token("ins", hyp)) + gen_spk = gen_ws[gen_idx] if gen_ws and gen_idx < len(gen_ws) else "" + if not cur_speaker: + cur_speaker = gen_spk + pred_tokens.append(Token("ins", hyp, speaker=gen_spk)) cur_pred.append(hyp) + gen_idx += 1 _flush() + seen: list[str] = [] + for t in ref_tokens: + if t.speaker and t.speaker not in seen: + seen.append(t.speaker) + for t in pred_tokens: + if t.speaker and t.speaker not in seen: + seen.append(t.speaker) + samples.append(Sample( file_stem=file_stem, example_num=int(ex.get("example", len(samples) + 1)), ref_tokens=tuple(ref_tokens), pred_tokens=tuple(pred_tokens), diffs=tuple(diffs), + speakers=tuple(seen), )) return samples def metrics_from_bewer(envelope: dict) -> dict: - """Metrics dict from a bewer envelope (wer/cer/mtr + normalization).""" + """Metrics dict from a bewer envelope (wer/cer/mtr + normalization). + + When the envelope carries ``per_speaker`` (from ``run_bewer_diarized``), + each speaker's metrics and word counts are included. + """ metrics = envelope.get("metrics") or {} settings = envelope.get("settings") or {} - return { + result = { "wer": metrics.get("wer"), "cer": metrics.get("cer"), "mtr": metrics.get("mtr"), "normalization": settings.get("normalization", True), } + per_speaker = envelope.get("per_speaker") or {} + if per_speaker: + result["per_speaker"] = { + label: { + "wer": sp["metrics"].get("wer"), + "cer": sp["metrics"].get("cer"), + "mtr": sp["metrics"].get("mtr"), + "ref_words": sp.get("ref_words", 0), + "gen_words": sp.get("gen_words", 0), + } + for label, sp in per_speaker.items() + } + diar_acc = envelope.get("diarization_accuracy") + if diar_acc: + result["diarization_accuracy"] = diar_acc.get("accuracy") + result["diarization_matched"] = diar_acc.get("matched") + result["diarization_total"] = diar_acc.get("total") + return result def reference_corpus(samples: list[Sample]) -> str: diff --git a/web/app.py b/web/app.py index dad7534..95ec203 100644 --- a/web/app.py +++ b/web/app.py @@ -38,6 +38,13 @@ from tympany import bewer_eval, corti from tympany.classify import classify_samples +from tympany.diarize import ( + Turn, + group_segments_by_turn, + is_diarized, + parse_corti_transcript_json, + split_into_turns, +) from tympany.llm import detect_provider from tympany.parser import from_bewer, metrics_from_bewer, reference_corpus, samples_to_payload from web import history, report_render, terms @@ -216,6 +223,7 @@ def _results_context( llm_notice: Optional[str] = None, ) -> dict: """Template context for results.html, shared by analyze / authored / history.""" + metrics = history.metrics_view(record) return { "user": _current_user(request), "nav_active": "analyze", @@ -226,9 +234,16 @@ def _results_context( "llm_notice": llm_notice, "analysis_id": record.get("id", ""), "download_base": history.download_base(record), - "metrics": history.metrics_view(record), + "metrics": metrics, "can_rerun": history.can_rerun(record), "is_authored": history.is_authored(record), + "has_speakers": any(e.get("speaker") for e in record.get("edits", [])), + "has_per_speaker": bool((record.get("original_metrics") or {}).get("per_speaker")), + "has_diarization_accuracy": metrics.get("has_diarization_accuracy"), + "speakers": sorted(set( + e.get("speaker", "") for e in record.get("edits", []) + if e.get("speaker") + )), } @@ -406,6 +421,12 @@ class _ReportInputs: report_name: str normalize_on: bool input_csv: str + ref_word_speakers: Optional[list[list[str]]] = None + gen_word_speakers: Optional[list[list[str]]] = None + diarized: bool = False + ref_segments: Optional[list] = None + gen_segments: Optional[list] = None + per_speaker: bool = False def _form_samples(references: list[str], generateds: list[str]) -> list[dict]: @@ -426,6 +447,9 @@ async def _prepare_report( csv_file: Optional[UploadFile], normalization: str, use_llm: str, terms_mode: str, terms_text: str, terms_file: Optional[UploadFile], terms_saved: str, terms_save_as: str, + ref_file: Optional[UploadFile] = None, + gen_file: Optional[UploadFile] = None, + per_speaker: str = "off", ): """Validate inputs and resolve medical terms for the create routes. @@ -441,7 +465,13 @@ async def _prepare_report( "terms_saved": terms_saved, "terms_save_as": terms_save_as, } try: - rows = await _read_rows(input_mode, reference, generated, csv_file) + if input_mode == "corti": + rows, ref_ws, gen_ws, diarized, ref_segs, gen_segs = await _read_corti_rows(ref_file, gen_file) + else: + rows = await _read_rows(input_mode, reference, generated, csv_file) + ref_ws = gen_ws = None + diarized = False + ref_segs = gen_segs = None terms_content = await _resolve_terms( email, terms_mode, terms_text, terms_file, terms_saved, terms_save_as ) @@ -454,15 +484,31 @@ async def _prepare_report( report_name=_safe_report_name(name), normalize_on=normalization == "on", input_csv=_rows_to_csv(rows), + ref_word_speakers=ref_ws, + gen_word_speakers=gen_ws, + diarized=diarized, + ref_segments=ref_segs if ref_segs else None, + gen_segments=gen_segs if gen_segs else None, + per_speaker=per_speaker == "on" and diarized, ) async def _run_bewer(prep: _ReportInputs) -> dict: """Evaluate prep's rows with bewer (off the event loop), returning the envelope.""" terms = list(_split_terms(prep.terms_content)) if prep.terms_content else None + if prep.per_speaker and prep.ref_segments and prep.gen_segments: + all_ref = [s for group in prep.ref_segments for s in group] + all_gen = [s for group in prep.gen_segments for s in group] + return await asyncio.to_thread( + bewer_eval.run_bewer_diarized, + all_ref, all_gen, + normalization=prep.normalize_on, medical_terms=terms, + ) return await asyncio.to_thread( bewer_eval.run_bewer, prep.rows, normalization=prep.normalize_on, medical_terms=terms, + ref_word_speakers=prep.ref_word_speakers, + gen_word_speakers=prep.gen_word_speakers, ) @@ -481,6 +527,9 @@ async def reports_generate( terms_file: Optional[UploadFile] = File(default=None), terms_saved: str = Form(default=""), terms_save_as: str = Form(default=""), + ref_file: Optional[UploadFile] = File(default=None), + gen_file: Optional[UploadFile] = File(default=None), + per_speaker: str = Form(default="off"), ): """Generate a BeWER report only (no analysis); stay on the page to iterate.""" redirect = _require_login(request) @@ -492,6 +541,7 @@ async def reports_generate( generated=generated, csv_file=csv_file, normalization=normalization, use_llm=use_llm, terms_mode=terms_mode, terms_text=terms_text, terms_file=terms_file, terms_saved=terms_saved, terms_save_as=terms_save_as, + ref_file=ref_file, gen_file=gen_file, per_speaker=per_speaker, ) if not isinstance(prep, _ReportInputs): return prep # a rendered error form @@ -535,6 +585,9 @@ async def reports_new_submit( terms_file: Optional[UploadFile] = File(default=None), terms_saved: str = Form(default=""), terms_save_as: str = Form(default=""), + ref_file: Optional[UploadFile] = File(default=None), + gen_file: Optional[UploadFile] = File(default=None), + per_speaker: str = Form(default="off"), ): """Generate a BeWER report and run the full Tympany analysis on it.""" redirect = _require_login(request) @@ -546,6 +599,7 @@ async def reports_new_submit( generated=generated, csv_file=csv_file, normalization=normalization, use_llm=use_llm, terms_mode=terms_mode, terms_text=terms_text, terms_file=terms_file, terms_saved=terms_saved, terms_save_as=terms_save_as, + ref_file=ref_file, gen_file=gen_file, per_speaker=per_speaker, ) if not isinstance(prep, _ReportInputs): return prep # a rendered error form @@ -562,10 +616,17 @@ async def reports_new_submit( error=f"BeWER could not generate a report: {exc}", ) - samples = from_bewer(envelope, file_stem=prep.report_name) + samples = from_bewer( + envelope, file_stem=prep.report_name, + ref_word_speakers=prep.ref_word_speakers, + gen_word_speakers=prep.gen_word_speakers, + ) llm_outcome: dict = {} edits = history.normalize_edits( - await asyncio.to_thread(classify_samples, samples, llm_provider, llm_outcome) + await asyncio.to_thread( + classify_samples, samples, llm_provider, llm_outcome, + ref_segments=prep.ref_segments, gen_segments=prep.gen_segments, + ) ) llm_used = bool(llm_provider) and not _llm_pass_failed(llm_outcome) @@ -578,10 +639,11 @@ async def reports_new_submit( source = { "kind": "authored", - "input": "csv" if input_mode == "csv" else "paste", + "input": "corti" if input_mode == "corti" else ("csv" if input_mode == "csv" else "paste"), "rows": len(prep.rows), "settings": {"normalization": prep.normalize_on}, "medical_terms": bool(prep.terms_content), + "diarized": prep.diarized, } analysis_id = history.save_analysis( email, filename, llm_used, edits, @@ -1012,6 +1074,35 @@ async def history_rerun(request: Request, analysis_id: str): updated = dict(envelope["metrics"]) updated["examples"] = len(rows) + + # Per-speaker re-run: reconstruct corrected text per speaker and re-evaluate. + orig_metrics = record.get("original_metrics") or {} + if orig_metrics.get("per_speaker"): + speaker_rows = bewer_eval.build_rows_per_speaker( + record.get("samples", []), record.get("edits", []) + ) + per_speaker_metrics: dict[str, dict] = {} + for spk, spk_rows in speaker_rows.items(): + if not spk_rows: + continue + try: + spk_env = await asyncio.to_thread( + bewer_eval.run_bewer, spk_rows, + normalization=settings.get("normalization", True), + medical_terms=terms, + ) + except bewer_eval.BewerError: + continue + per_speaker_metrics[spk] = { + "wer": spk_env["metrics"].get("wer"), + "cer": spk_env["metrics"].get("cer"), + "mtr": spk_env["metrics"].get("mtr"), + "ref_words": sum(len(r[0].split()) for r in spk_rows), + "gen_words": sum(len(r[1].split()) for r in spk_rows), + } + if per_speaker_metrics: + updated["per_speaker"] = per_speaker_metrics + history.set_bewer_rerun(email, analysis_id, updated, json.dumps(envelope)) return JSONResponse({"ok": True, "metrics": history.metrics_view( history.load_analysis(email, analysis_id) @@ -1294,6 +1385,89 @@ async def _read_rows( return rows +async def _read_corti_rows( + ref_file: Optional[UploadFile], + gen_file: Optional[UploadFile], +) -> tuple[ + list[tuple[str, str]], + Optional[list[list[str]]], + Optional[list[list[str]]], + bool, + Optional[list], + Optional[list], +]: + """Parse Corti diarized transcript JSON uploads into per-turn (ref, gen) rows. + + Returns ``(rows, ref_word_speakers, gen_word_speakers, diarized, + ref_segment_groups, gen_segment_groups)``. Each Corti transcript is split + into turns (consecutive same-speaker segments); turns are paired by + position to produce one bewer example per turn. ``ref_segment_groups`` + and ``gen_segment_groups`` are ``list[list[SpeakerSegment]]`` — one list + of raw segments per turn — so ``classify_samples`` can run diarization + error detection per turn rather than across the entire transcript. + """ + if ref_file is None or not ref_file.filename: + raise ValueError("Upload a reference transcript JSON file.") + if gen_file is None or not gen_file.filename: + raise ValueError("Upload a generated transcript JSON file.") + if not ref_file.filename.lower().endswith(".json"): + raise ValueError("Reference file must be a .json Corti transcript.") + if not gen_file.filename.lower().endswith(".json"): + raise ValueError("Generated file must be a .json Corti transcript.") + + ref_raw = await ref_file.read() + gen_raw = await gen_file.read() + for raw, label in ((ref_raw, "Reference"), (gen_raw, "Generated")): + if len(raw) > MAX_UPLOAD_BYTES: + raise ValueError(f"{label} file is too large (limit {MAX_UPLOAD_BYTES // (1024 * 1024)} MB).") + + try: + ref_segments = parse_corti_transcript_json(_decode_upload(ref_raw)) + except ValueError as exc: + raise ValueError(f"Reference transcript: {exc}") + try: + gen_segments = parse_corti_transcript_json(_decode_upload(gen_raw)) + except ValueError as exc: + raise ValueError(f"Generated transcript: {exc}") + + if not ref_segments: + raise ValueError("No transcript segments found in the reference file.") + if not gen_segments: + raise ValueError("No transcript segments found in the generated file.") + + ref_turns = split_into_turns(ref_segments) + gen_turns = split_into_turns(gen_segments) + + if not any(t.text for t in ref_turns): + raise ValueError("Reference transcript segments contain no text.") + if not any(t.text for t in gen_turns): + raise ValueError("Generated transcript segments contain no text.") + + n = max(len(ref_turns), len(gen_turns)) + rows: list[tuple[str, str]] = [] + ref_ws: list[list[str]] = [] + gen_ws: list[list[str]] = [] + for i in range(n): + rt = ref_turns[i] if i < len(ref_turns) else Turn("", "") + gt = gen_turns[i] if i < len(gen_turns) else Turn("", "") + rows.append((rt.text, gt.text)) + ref_ws.append([rt.speaker] * len(rt.text.split()) if rt.text else []) + gen_ws.append([gt.speaker] * len(gt.text.split()) if gt.text else []) + + diarized = is_diarized(ref_segments) or is_diarized(gen_segments) + + ref_seg_groups = group_segments_by_turn(ref_segments) + gen_seg_groups = group_segments_by_turn(gen_segments) + while len(ref_seg_groups) < n: + ref_seg_groups.append([]) + while len(gen_seg_groups) < n: + gen_seg_groups.append([]) + + if len(rows) > MAX_ROWS: + raise ValueError(f"Too many examples ({len(rows)}); the limit is {MAX_ROWS}.") + return rows, ref_ws, gen_ws, diarized, ref_seg_groups, gen_seg_groups + + async def _resolve_terms( email: str, terms_mode: str, diff --git a/web/history.py b/web/history.py index 482e4db..c1530e3 100644 --- a/web/history.py +++ b/web/history.py @@ -31,6 +31,7 @@ "context_dependent", "formatting_error", "replacement_candidate", + "diarization_error", ] _RISK_LEVELS = ["high", "medium", "low"] @@ -98,6 +99,7 @@ def normalize_edits(rows: list[dict]) -> list[dict]: out.append({ "file": row.get("file", ""), "example": row.get("example", ""), + "speaker": row.get("speaker", ""), "ref": row.get("ref", ""), "gen": row.get("gen", ""), "op": row.get("op", ""), @@ -257,7 +259,7 @@ def delete_analysis(email: str, analysis_id: str) -> bool: # Column order matches the client-side exporter in results.html so server and # browser downloads produce byte-identical CSVs from the same edit state. CSV_HEADER = [ - "file", "example", "ref", "gen", "op", + "file", "example", "speaker", "ref", "gen", "op", "classification", "risk_level", "error_description", "excluded", "flagged", "detail", ] @@ -273,6 +275,7 @@ def to_csv(record: dict) -> str: writer.writerow([ e.get("file", ""), e.get("example", ""), + e.get("speaker", ""), e.get("ref", ""), e.get("gen", ""), e.get("op", ""), @@ -293,7 +296,7 @@ def to_csv(record: dict) -> str: # Like CSV_HEADER but adds the full example reference / generated text so a # reviewer has the context surrounding each flagged error, not just the diff. FLAGS_CSV_HEADER = [ - "file", "example", "ref", "gen", "op", + "file", "example", "speaker", "ref", "gen", "op", "classification", "risk_level", "error_description", "detail", "ref_context", "gen_context", ] @@ -326,6 +329,7 @@ def flags_to_csv(record: dict) -> str: writer.writerow([ e.get("file", ""), e.get("example", ""), + e.get("speaker", ""), e.get("ref", ""), e.get("gen", ""), e.get("op", ""), @@ -610,6 +614,7 @@ def set_bewer_rerun( "wer": updated_metrics.get("wer"), "cer": updated_metrics.get("cer"), "mtr": updated_metrics.get("mtr"), + "per_speaker": updated_metrics.get("per_speaker"), }, "examples": updated_metrics.get("examples"), "run_at": _now_iso(), @@ -644,6 +649,21 @@ def metrics_view(record: dict) -> dict: orig_mtr, upd_mtr = _pct(original.get("mtr")), _pct(updated.get("mtr")) mtr_improved = orig_mtr is not None and upd_mtr is not None and upd_mtr > orig_mtr + # Per-speaker metrics (Phase 2): merge original + updated into one view. + orig_ps = original.get("per_speaker") or {} + upd_ps = updated.get("per_speaker") or {} + per_speaker: dict[str, dict] = {} + for spk in list(orig_ps.keys()) + [s for s in upd_ps if s not in orig_ps]: + o = orig_ps.get(spk, {}) + u = upd_ps.get(spk, {}) + per_speaker[spk] = { + "original_wer": o.get("wer"), + "original_cer": o.get("cer"), + "updated_wer": u.get("wer"), + "updated_cer": u.get("cer"), + "ref_words": o.get("ref_words", 0), + } + return { "original_wer": original.get("wer"), "original_cer": original.get("cer"), @@ -653,6 +673,12 @@ def metrics_view(record: dict) -> dict: "updated_mtr": updated.get("mtr"), "mtr_improved": mtr_improved, "has_rerun": bool(updated.get("wer") or updated.get("cer")), + "per_speaker": per_speaker, + "has_per_speaker": bool(per_speaker), + "diarization_accuracy": original.get("diarization_accuracy"), + "diarization_matched": original.get("diarization_matched"), + "diarization_total": original.get("diarization_total"), + "has_diarization_accuracy": original.get("diarization_accuracy") is not None, } diff --git a/web/report_render.py b/web/report_render.py index dde3c42..02c885f 100644 --- a/web/report_render.py +++ b/web/report_render.py @@ -209,6 +209,47 @@ def _example_lines(ex: dict, fallback_terms: list[tuple[str, ...]]) -> list[dict ] +def _split_ops_by_speaker( + ex: dict, +) -> list[tuple[str, list[dict]]]: + """Split an example's alignment ops into per-speaker groups. + + Walks the ops tracking ref/gen token indices, assigns a speaker label + to each op (from ``ref_speakers`` for MATCH/SUBSTITUTE/DELETE, + ``hyp_speakers`` for INSERT), and groups consecutive ops with the same + speaker. Returns a list of ``(speaker_label, ops_list)`` tuples. + + Ops with no speaker label (non-diarized) are grouped under an empty + string, so the caller can detect "no split needed" and render a single + block. + """ + ops = ex.get("ops") or [] + ref_spk = ex.get("ref_speakers") or [] + hyp_spk = ex.get("hyp_speakers") or [] + if not ref_spk and not hyp_spk: + return [("", list(ops))] + + groups: list[tuple[str, list[dict]]] = [] + ref_idx = 0 + gen_idx = 0 + for op in ops: + op_type = (op.get("type") or "").upper() + if op_type == "INSERT": + spk = hyp_spk[gen_idx] if gen_idx < len(hyp_spk) else "" + gen_idx += 1 + else: + spk = ref_spk[ref_idx] if ref_idx < len(ref_spk) else "" + if op_type in ("MATCH", "SUBSTITUTE"): + gen_idx += 1 + ref_idx += 1 + + if groups and groups[-1][0] == spk: + groups[-1][1].append(op) + else: + groups.append((spk, [op])) + return groups + + def _comma(n: int) -> str: return f"{n:,}" @@ -238,10 +279,30 @@ def canal_view(envelope: dict) -> dict: if op_type in ("MATCH", "SUBSTITUTE", "INSERT"): gen_words += 1 gen_chars += len(hyp) - examples.append({ - "example": ex.get("example"), - "lines": _example_lines(ex, fallback_terms), - }) + ref_spk = ex.get("ref_speakers") or [] + hyp_spk = ex.get("hyp_speakers") or [] + all_speakers: list[str] = [] + for s in ref_spk + hyp_spk: + if s and s not in all_speakers: + all_speakers.append(s) + + if all_speakers: + for spk, spk_ops in _split_ops_by_speaker(ex): + label = spk or all_speakers[0] + sub_ex = {"ops": spk_ops} + examples.append({ + "example": ex.get("example"), + "speaker": label, + "lines": _example_lines(sub_ex, fallback_terms), + "speakers": all_speakers, + }) + else: + examples.append({ + "example": ex.get("example"), + "speaker": "", + "lines": _example_lines(ex, fallback_terms), + "speakers": [], + }) mtr = metrics.get("mtr") summary = { @@ -255,11 +316,28 @@ def canal_view(envelope: dict) -> dict: "gen_chars": _comma(gen_chars), } + per_speaker = envelope.get("per_speaker") or {} + per_speaker_rows = [ + { + "label": label, + "wer": sp.get("metrics", {}).get("wer") or "—", + "cer": sp.get("metrics", {}).get("cer") or "—", + "mtr": sp.get("metrics", {}).get("mtr"), + "ref_words": _comma(sp.get("ref_words", 0)), + "gen_words": _comma(sp.get("gen_words", 0)), + } + for label, sp in per_speaker.items() + ] + return { "summary": summary, "examples": examples, "show_mtr": bool(mtr), - # Show the Medical Term legend whenever a key-term list was supplied. + "show_speakers": any(ex.get("speakers") for ex in examples), + "show_per_speaker": bool(per_speaker_rows), + "per_speaker_rows": per_speaker_rows, + "diarization_accuracy": envelope.get("diarization_accuracy"), + "show_diarization_accuracy": envelope.get("diarization_accuracy") is not None, "medical_terms": bool(key_terms), "normalization": bool(settings.get("normalization", True)), } diff --git a/web/templates/bewer_report.html b/web/templates/bewer_report.html index eb2169d..a3371dc 100644 --- a/web/templates/bewer_report.html +++ b/web/templates/bewer_report.html @@ -213,6 +213,7 @@

{{ report_title or "BeWER report" }}

Word Error Rate{{ summary.wer }} Character Error Rate{{ summary.cer }} {% if show_mtr %}Medical Term Recall{{ summary.mtr }}{% endif %} + {% if show_diarization_accuracy %}Diarization Accuracy{{ diarization_accuracy.accuracy }} ({{ diarization_accuracy.matched }}/{{ diarization_accuracy.total }} turns){% endif %} @@ -228,6 +229,30 @@

{{ report_title or "BeWER report" }}

+ {% if show_per_speaker %} +
+ + + + + + + {% if show_mtr %}{% endif %} + + + {% for row in per_speaker_rows %} + + + + + {% if show_mtr %}{% endif %} + + + {% endfor %} +
Per-Speaker Metrics
SpeakerWERCERMTRRef words
{{ row.label }}{{ row.wer }}{{ row.cer }}{{ row.mtr or "—" }}{{ row.ref_words }}
+
+ {% endif %} +
Examples
+ +
+
+
+ + +
+
+ + +
+
+
Upload Corti transcript JSON (streams or REST /transcripts format). Speaker labels are parsed from speakerId and channel; text is flattened for evaluation and each edit row is tagged with its speaker.
+
+ {% set tmode = f.terms_mode or "none" %}
@@ -157,6 +173,14 @@
+
+ + +
+
{% if llm_provider %} @@ -232,12 +256,16 @@ const tabs = document.querySelectorAll(".mode-tab"); const panePaste = document.getElementById("pane-paste"); const paneCsv = document.getElementById("pane-csv"); + const paneCorti = document.getElementById("pane-corti"); function setMode(mode) { hidden.value = mode; tabs.forEach(t => t.classList.toggle("active", t.dataset.mode === mode)); - panePaste.style.display = mode === "csv" ? "none" : ""; + panePaste.style.display = mode === "paste" ? "" : "none"; paneCsv.style.display = mode === "csv" ? "" : "none"; + paneCorti.style.display = mode === "corti" ? "" : "none"; + const psRow = document.getElementById("per-speaker-row"); + if (psRow) psRow.style.display = mode === "corti" ? "flex" : "none"; } tabs.forEach(t => t.addEventListener("click", () => setMode(t.dataset.mode))); })(); diff --git a/web/templates/results.html b/web/templates/results.html index 40cac44..3f636a1 100644 --- a/web/templates/results.html +++ b/web/templates/results.html @@ -81,6 +81,7 @@ .cls-select.cls-replacement_candidate { color: hsl(var(--variant-info-text)); } .cls-select.cls-context_dependent { color: hsl(var(--variant-warning-text)); } .cls-select.cls-misrecognition { color: hsl(var(--variant-error-text)); } + .cls-select.cls-diarization_error { color: hsl(var(--primary)); } /* Editable description input */ .desc-input { @@ -360,6 +361,50 @@
+ + {% if has_per_speaker %} +
+
Per-speaker evaluation
+ + + + + + + + + + + + + {% for spk, m in metrics.per_speaker.items() %} + + + + + + + + + {% endfor %} + +
SpeakerWER (Original)WER (Updated)CER (Original)CER (Updated)Ref words
{{ spk }}{{ m.original_wer or "—" }}{{ m.updated_wer or "—" }}{{ m.original_cer or "—" }}{{ m.updated_cer or "—" }}{{ m.ref_words }}
+
+ {% endif %} + + + {% if has_diarization_accuracy %} +
+
Diarization Accuracy
+
+
+ {{ metrics.diarization_accuracy or "—" }} + {{ metrics.diarization_matched }} / {{ metrics.diarization_total }} turns matched +
+
+
+ {% endif %} +
@@ -390,6 +435,7 @@ Context-dep. Formatting Replacement + Diariz. @@ -399,6 +445,7 @@ — — — + — Excluded @@ -406,6 +453,7 @@ — — — + — @@ -529,6 +577,15 @@
+ {% if has_speakers %} + + + {% endif %} @@ -668,6 +733,7 @@ replacement_candidate: 'low', context_dependent: 'medium', misrecognition: 'high', + diarization_error: 'medium', }; const ANALYSIS_ID = document.querySelector('#edits-table')?.dataset.analysisId || ''; @@ -731,6 +797,7 @@ return [...document.querySelectorAll('#edits-table tbody tr')].map(row => ({ file: row.dataset.file || '', example: row.dataset.example || '', + speaker: row.dataset.speaker || '', ref: row.dataset.ref || '', gen: row.dataset.gen || '', op: row.dataset.op || '', @@ -786,7 +853,7 @@ document.getElementById('stat-total').textContent = active.length; // By classification table - ['misrecognition','context_dependent','formatting_error','replacement_candidate'].forEach(cls => { + ['misrecognition','context_dependent','formatting_error','replacement_candidate','diarization_error'].forEach(cls => { const tot = rows.filter(r => r.dataset.cls === cls).length; const exc = excluded.filter(r => r.dataset.cls === cls).length; document.getElementById('bt-total-' + cls).textContent = tot; @@ -877,6 +944,7 @@ const cls = document.getElementById('filter-cls').value; const risk = document.getElementById('filter-risk').value; const excludedMode = document.getElementById('filter-excluded').value; + const speaker = document.getElementById('filter-speaker')?.value || ''; const rows = document.querySelectorAll('#edits-table tbody tr'); let visible = 0; rows.forEach(row => { @@ -889,6 +957,7 @@ (activeSamples.size === 0 || activeSamples.has(row.dataset.example)) && (!cls || row.dataset.cls === cls) && (!risk || row.dataset.risk === risk) && + (!speaker || row.dataset.speaker === speaker) && matchesExcluded; row.classList.toggle('hidden', !show); if (show) visible++; @@ -899,7 +968,7 @@ // ── CSV download (client-side, reflects all edits) ─────────────────────────── function downloadCSV() { const rows = [...document.querySelectorAll('#edits-table tbody tr')]; - const header = ['file','example','ref','gen','op','classification','risk_level','error_description','excluded','flagged','detail']; + const header = ['file','example','speaker','ref','gen','op','classification','risk_level','error_description','excluded','flagged','detail']; const lines = [header.join(',')]; rows.forEach(row => { @@ -912,6 +981,7 @@ const cells = [ csvEsc(row.dataset.file || ''), csvEsc(row.dataset.example || ''), + csvEsc(row.dataset.speaker || ''), csvEsc(row.dataset.ref || ''), csvEsc(row.dataset.gen || ''), csvEsc(row.dataset.op || ''), @@ -999,6 +1069,15 @@ if (link && m.has_rerun) link.style.visibility = 'visible'; const viewLink = document.getElementById('view-rerun-link'); if (viewLink && m.has_rerun) viewLink.style.display = ''; + // Per-speaker updated metrics + if (m.per_speaker) { + for (const [spk, vals] of Object.entries(m.per_speaker)) { + const werCell = document.querySelector(`.spk-upd-wer[data-speaker="${spk}"]`); + const cerCell = document.querySelector(`.spk-upd-cer[data-speaker="${spk}"]`); + if (werCell && vals.updated_wer) werCell.textContent = vals.updated_wer; + if (cerCell && vals.updated_cer) cerCell.textContent = vals.updated_cer; + } + } } // ── Init ─────────────────────────────────────────────────────────────────────