Bugfixes in the LLM retry loop
This commit is contained in:
@@ -195,15 +195,13 @@ def _run_correction_stage(
|
|||||||
)
|
)
|
||||||
working = application_result.transcript
|
working = application_result.transcript
|
||||||
|
|
||||||
for correction_id in application_result.applied_ids:
|
next_retry_skips: Dict[int, SkippedCorrection] = {}
|
||||||
unresolved_retry_skips.pop(correction_id, None)
|
|
||||||
for correction_id in application_result.ignored_ids:
|
|
||||||
unresolved_retry_skips.pop(correction_id, None)
|
|
||||||
for skipped in application_result.skipped:
|
for skipped in application_result.skipped:
|
||||||
if _is_retryable_skip(skipped, working):
|
if _is_retryable_skip(skipped, working):
|
||||||
unresolved_retry_skips[skipped.id] = skipped
|
next_retry_skips[skipped.id] = skipped
|
||||||
else:
|
else:
|
||||||
final_nonretry_skips.append(skipped)
|
final_nonretry_skips.append(skipped)
|
||||||
|
unresolved_retry_skips = next_retry_skips
|
||||||
|
|
||||||
pass_summaries.append(
|
pass_summaries.append(
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -310,6 +310,110 @@ def test_pipeline_retries_skipped_segment_and_cleans_work_dir_when_fixed(tmp_pat
|
|||||||
assert list((tmp_path / "work").iterdir()) == []
|
assert list((tmp_path / "work").iterdir()) == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_empty_glossary_retry_response_stops_retrying_segment(tmp_path):
|
||||||
|
first_pass = CorrectionCandidate(
|
||||||
|
id=1,
|
||||||
|
original_text="Contia",
|
||||||
|
corrected_text="Chauntea",
|
||||||
|
confidence=0.95,
|
||||||
|
)
|
||||||
|
fake_client = FakeLLMClient(
|
||||||
|
[
|
||||||
|
CorrectionSet(corrections=[first_pass]),
|
||||||
|
CorrectionSet(corrections=[]),
|
||||||
|
CorrectionSet(corrections=[]),
|
||||||
|
CorrectionSet(corrections=[]),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
revised = process_transcript(
|
||||||
|
_transcript(),
|
||||||
|
_glossary(),
|
||||||
|
_config(tmp_path, glossary_max_llm_passes=3),
|
||||||
|
llm_client=fake_client,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert fake_client.calls == 3
|
||||||
|
assert "readability corrections" in fake_client.messages[2][1]["content"]
|
||||||
|
assert revised[0].text == "I ask Chontia."
|
||||||
|
assert list((tmp_path / "work").iterdir()) == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_empty_grammar_retry_response_stops_retrying_segment(tmp_path):
|
||||||
|
transcript = parse_source_transcript_json(
|
||||||
|
"""
|
||||||
|
[
|
||||||
|
{"speaker": "Eric", "start": 0.0, "end": 1.0, "text": "there and there."}
|
||||||
|
]
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
repeated_span = CorrectionCandidate(
|
||||||
|
id=1,
|
||||||
|
original_text="there",
|
||||||
|
corrected_text="their",
|
||||||
|
confidence=0.95,
|
||||||
|
)
|
||||||
|
fake_client = FakeLLMClient(
|
||||||
|
[
|
||||||
|
CorrectionSet(corrections=[]),
|
||||||
|
CorrectionSet(corrections=[repeated_span]),
|
||||||
|
CorrectionSet(corrections=[]),
|
||||||
|
CorrectionSet(corrections=[]),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
revised = process_transcript(
|
||||||
|
transcript,
|
||||||
|
_glossary(),
|
||||||
|
_config(tmp_path, grammar_max_llm_passes=3),
|
||||||
|
llm_client=fake_client,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert fake_client.calls == 3
|
||||||
|
assert revised[0].text == "there and there."
|
||||||
|
assert list((tmp_path / "work").iterdir()) == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_retry_pass_with_new_skip_schedules_following_pass(tmp_path):
|
||||||
|
first_pass = CorrectionCandidate(
|
||||||
|
id=1,
|
||||||
|
original_text="Contia",
|
||||||
|
corrected_text="Chauntea",
|
||||||
|
confidence=0.95,
|
||||||
|
)
|
||||||
|
second_pass = CorrectionCandidate(
|
||||||
|
id=1,
|
||||||
|
original_text="Chantia",
|
||||||
|
corrected_text="Chauntea",
|
||||||
|
confidence=0.95,
|
||||||
|
)
|
||||||
|
third_pass = CorrectionCandidate(
|
||||||
|
id=1,
|
||||||
|
original_text="Chontia",
|
||||||
|
corrected_text="Chauntea",
|
||||||
|
confidence=0.95,
|
||||||
|
)
|
||||||
|
fake_client = FakeLLMClient(
|
||||||
|
[
|
||||||
|
CorrectionSet(corrections=[first_pass]),
|
||||||
|
CorrectionSet(corrections=[second_pass]),
|
||||||
|
CorrectionSet(corrections=[third_pass]),
|
||||||
|
CorrectionSet(corrections=[]),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
revised = process_transcript(
|
||||||
|
_transcript(),
|
||||||
|
_glossary(),
|
||||||
|
_config(tmp_path, glossary_max_llm_passes=3),
|
||||||
|
llm_client=fake_client,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert fake_client.calls == 4
|
||||||
|
assert revised[0].text == "I ask Chauntea."
|
||||||
|
assert list((tmp_path / "work").iterdir()) == []
|
||||||
|
|
||||||
|
|
||||||
def test_pipeline_retry_prompt_contains_only_valid_deduped_ids(tmp_path):
|
def test_pipeline_retry_prompt_contains_only_valid_deduped_ids(tmp_path):
|
||||||
first_bad = CorrectionCandidate(
|
first_bad = CorrectionCandidate(
|
||||||
id=1,
|
id=1,
|
||||||
@@ -358,10 +462,16 @@ def test_pipeline_writes_stage_metadata_for_unresolved_retries(tmp_path):
|
|||||||
corrected_text="Chauntea",
|
corrected_text="Chauntea",
|
||||||
confidence=0.95,
|
confidence=0.95,
|
||||||
)
|
)
|
||||||
|
second_pass = CorrectionCandidate(
|
||||||
|
id=1,
|
||||||
|
original_text="Chantia",
|
||||||
|
corrected_text="Chauntea",
|
||||||
|
confidence=0.95,
|
||||||
|
)
|
||||||
fake_client = FakeLLMClient(
|
fake_client = FakeLLMClient(
|
||||||
[
|
[
|
||||||
CorrectionSet(corrections=[first_pass]),
|
CorrectionSet(corrections=[first_pass]),
|
||||||
CorrectionSet(corrections=[]),
|
CorrectionSet(corrections=[second_pass]),
|
||||||
CorrectionSet(corrections=[]),
|
CorrectionSet(corrections=[]),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
@@ -392,6 +502,7 @@ def test_pipeline_writes_stage_metadata_for_unresolved_retries(tmp_path):
|
|||||||
assert [item["pass_number"] for item in metadata["stages"][0]["passes"]] == [1, 2]
|
assert [item["pass_number"] for item in metadata["stages"][0]["passes"]] == [1, 2]
|
||||||
assert metadata["stages"][0]["passes"][0]["retry_segment_count"] == 1
|
assert metadata["stages"][0]["passes"][0]["retry_segment_count"] == 1
|
||||||
assert metadata["stages"][0]["passes"][1]["retry_pass"] is True
|
assert metadata["stages"][0]["passes"][1]["retry_pass"] is True
|
||||||
|
assert metadata["stages"][0]["passes"][1]["retry_segment_count"] == 1
|
||||||
|
|
||||||
|
|
||||||
def test_grammar_stage_runs_after_glossary_and_sees_corrected_text(tmp_path):
|
def test_grammar_stage_runs_after_glossary_and_sees_corrected_text(tmp_path):
|
||||||
|
|||||||
Reference in New Issue
Block a user