Handle repetitive Semantic Consolidator output loops

This commit is contained in:
2026-08-09 14:33:46 +02:00
parent 3a5850430b
commit 58effcafc5
5 changed files with 462 additions and 11 deletions
@@ -6,6 +6,7 @@ from __future__ import annotations
import argparse
import copy
import json
import re
import sys
import time
from collections import Counter
@@ -34,6 +35,15 @@ OUTPUT_TOKEN_ESTIMATE_CHARS = 4
OUTPUT_GROUP_OVERHEAD_CHARS = 320
OUTPUT_SAFETY_MARGIN = 1.35
PROMPT_NAME = "consolidate_facts.md"
REPETITION_LOOP_MIN_CONSECUTIVE_GROUPS = 3
REPETITION_RETRY_INSTRUCTION = """
RETRY SAFETY INSTRUCTION:
Emit each semantic group exactly once. Before returning the JSON, check the
groups array for identical groups using canonical_text, source_item_ids and
merge_reason. If an identical group is already present, do not emit it again.
Finish the complete JSON object without repeating any group.
""".rstrip()
class ConsolidationValidationError(ValueError):
@@ -343,6 +353,169 @@ def parse_model_json(text: str) -> dict[str, Any]:
return parsed
def group_repetition_signature(group: dict[str, Any]) -> tuple[str, tuple[str, ...], str] | None:
canonical_text = group.get("canonical_text")
source_item_ids = group.get("source_item_ids")
merge_reason = group.get("merge_reason")
if (
not isinstance(canonical_text, str)
or not isinstance(source_item_ids, list)
or not all(isinstance(item_id, str) for item_id in source_item_ids)
or not isinstance(merge_reason, str)
):
return None
normalize = lambda value: re.sub(r"\s+", " ", value).strip()
return (
normalize(canonical_text),
tuple(source_item_ids),
normalize(merge_reason),
)
def completed_group_objects(text: str) -> list[dict[str, Any]]:
"""Extract complete JSON group objects even when the outer response is truncated."""
groups_key = re.search(r'"groups"\s*:\s*\[', text)
if groups_key is None:
return []
decoder = json.JSONDecoder()
groups: list[dict[str, Any]] = []
position = groups_key.end()
while position < len(text):
while position < len(text) and text[position] in " \t\r\n,":
position += 1
if position >= len(text) or text[position] == "]":
break
if text[position] != "{":
break
try:
value, position = decoder.raw_decode(text, position)
except json.JSONDecodeError:
break
if not isinstance(value, dict):
break
groups.append(value)
return groups
def detect_repetitive_group_loop(
text: str,
minimum_consecutive: int = REPETITION_LOOP_MIN_CONSECUTIVE_GROUPS,
) -> dict[str, Any]:
groups = completed_group_objects(text)
longest_run = 0
longest_signature: tuple[str, tuple[str, ...], str] | None = None
run_start = 0
previous: tuple[str, tuple[str, ...], str] | None = None
current_run = 0
for index, group in enumerate(groups):
signature = group_repetition_signature(group)
if signature is not None and signature == previous:
current_run += 1
else:
current_run = 1 if signature is not None else 0
run_start = index
if current_run > longest_run:
longest_run = current_run
longest_signature = signature
longest_run_start = run_start
previous = signature
detected = longest_signature is not None and longest_run >= minimum_consecutive
signature_data = None
if longest_signature is not None:
signature_data = {
"canonical_text": longest_signature[0],
"source_item_ids": list(longest_signature[1]),
"merge_reason": longest_signature[2],
}
return {
"detected": detected,
"threshold": minimum_consecutive,
"completed_group_count": len(groups),
"longest_consecutive_run": longest_run,
"run_start_group_index": longest_run_start if longest_run else None,
"signature": signature_data,
}
def call_ollama_with_repetition_retry(
*,
endpoint: str,
model: str,
prompt: str,
timeout: int,
num_ctx: int,
num_predict: int,
think: bool,
progress_interval: int,
output_dir: Path,
call_fn: Any = None,
) -> tuple[str, dict[str, Any], float, list[dict[str, Any]]]:
"""Make one call, retrying once only for invalid JSON with a detected group loop."""
if call_fn is None:
call_fn = call_ollama
output_dir.mkdir(parents=True, exist_ok=True)
attempts: list[dict[str, Any]] = []
total_runtime = 0.0
for attempt_number in (1, 2):
attempt_prompt = (
prompt
if attempt_number == 1
else f"{prompt}\n{REPETITION_RETRY_INSTRUCTION}\n"
)
raw_text, response_data, runtime = call_fn(
endpoint=endpoint,
model=model,
prompt=attempt_prompt,
timeout=timeout,
num_ctx=num_ctx,
num_predict=num_predict,
think=think,
progress_interval=progress_interval,
)
total_runtime += runtime
suffix = "" if attempt_number == 1 else "_retry"
(output_dir / f"raw_model_response{suffix}.txt").write_text(
raw_text + "\n", encoding="utf-8"
)
write_json(output_dir / f"raw_ollama_response{suffix}.json", response_data)
repetition = detect_repetitive_group_loop(raw_text)
attempt_metadata = {
"attempt": attempt_number,
"retry_instruction_added": attempt_number == 2,
"runtime_seconds": runtime,
"raw_response_chars": len(raw_text),
"raw_response_bytes": len(raw_text.encode("utf-8")),
"resolved_num_predict": num_predict,
"eval_count": response_data.get("eval_count"),
"done_reason": response_data.get("done_reason"),
"repetition": repetition,
}
attempts.append(attempt_metadata)
try:
parse_model_json(raw_text)
except ConsolidationValidationError as exc:
attempt_metadata["parse_error"] = str(exc)
write_json(
output_dir / "repetition_retry_metadata.json",
{"bug": "BUG-013", "attempts": attempts},
)
if attempt_number == 1 and repetition["detected"]:
continue
raise
write_json(
output_dir / "repetition_retry_metadata.json",
{"bug": "BUG-013", "attempts": attempts},
)
return raw_text, response_data, total_runtime, attempts
raise AssertionError("repetition retry loop exhausted unexpectedly")
def validate_model_groups(
model_output: dict[str, Any],
expected_fact_ids: set[str],
@@ -626,6 +799,7 @@ def write_report(
groups: list[dict[str, Any]],
output_path: Path,
repair_changes: list[dict[str, Any]] | None = None,
llm_call_count: int = 1,
) -> None:
merged = [group for group in groups if len(group["source_item_ids"]) > 1]
singletons = [group for group in groups if len(group["source_item_ids"]) == 1]
@@ -634,7 +808,7 @@ def write_report(
"",
"- Scope: facts only",
f"- Model: `{model}`",
"- LLM call count: 1",
f"- LLM call count: {llm_call_count}",
f"- Runtime: {runtime:.3f} seconds",
f"- Fact item count: {fact_count}",
f"- Prompt characters: {prompt_chars}",
@@ -693,10 +867,10 @@ def main() -> int:
print(f"Estimated prompt size chars: {prompt_chars}")
print(f"Estimated prompt tokens: {prompt_token_estimate}")
print(f"Resolved num_predict: {num_predict}")
print("Expected LLM call count: 1")
print("Expected LLM call count: 1; at most 2 only after detected BUG-013 loop")
print("Expected runtime: 5-10 minutes on current local benchmark basis")
raw_text, _response_data, runtime = call_ollama(
raw_text, _response_data, runtime, attempts = call_ollama_with_repetition_retry(
endpoint=args.endpoint,
model=args.model,
prompt=prompt,
@@ -705,8 +879,8 @@ def main() -> int:
num_predict=num_predict,
think=args.think,
progress_interval=args.progress_interval,
output_dir=args.output_dir,
)
raw_response_path.write_text(raw_text + "\n", encoding="utf-8")
model_output = parse_model_json(raw_text)
model_output, repair_changes = repair_model_group_coverage(model_output, facts)
if repair_changes:
@@ -736,6 +910,7 @@ def main() -> int:
groups=groups,
output_path=output_path,
repair_changes=repair_changes,
llm_call_count=len(attempts),
)
except requests.ConnectionError as exc:
print(f"Error: Ollama is not reachable at {args.endpoint}: {exc}", file=sys.stderr)