Handle repetitive Semantic Consolidator output loops
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user