#!/usr/bin/env python3
"""Validate synthetic signature-to-background dataset JSONL files."""
from __future__ import annotations
import json
import re
from collections import Counter
from pathlib import Path
ROOT = Path(__file__).resolve().parents[2]
PROCESSED = ROOT / "dataset" / "processed"
CATALOG_PATH = ROOT / "dataset" / "config" / "process_catalog.v1.json"
TASK_TYPE = "signature_to_backgrounds"
SPLITS = {"train", "val", "test"}
SFT_TARGET_TYPES = {
"dominant_irreducible": "irreducible",
"dominant_reducible": "reducible",
}
SFT_HEADERS = ["dominant:", "irreducible:", "reducible:"]
TRAINABLE_EVIDENCE_TRACE_RE = re.compile(
r"(?im)(evidence trace|^\s*evidence\s*:|^\s*citations?\s*:|^\s*sources?\s*:|^\s*references?\s*:|source_file|claim_supported)"
)
THINK_BLOCK_RE = re.compile(r"(?is)(.*?)")
ANSWER_BLOCK_RE = re.compile(r"(?is)\s*(\{.*?\})\s*")
ANSWER_TEXT_BLOCK_RE = re.compile(r"(?is)\s*(.*?)\s*")
def load_catalog() -> tuple[str, set[str]]:
with CATALOG_PATH.open() as handle:
catalog = json.load(handle)
version = str(catalog["version"])
process_ids = {str(process["id"]) for process in catalog["processes"]}
if len(process_ids) != len(catalog["processes"]):
raise ValueError(f"{CATALOG_PATH}: duplicate process ids")
return version, process_ids
PROCESS_CATALOG_VERSION, PROCESS_IDS = load_catalog()
def read_jsonl(path: Path) -> list[dict]:
rows: list[dict] = []
with path.open() as handle:
for line_no, line in enumerate(handle, 1):
if not line.strip():
continue
try:
rows.append(json.loads(line))
except json.JSONDecodeError as exc:
raise ValueError(f"{path}:{line_no}: invalid JSON: {exc}") from exc
return rows
def section_items(text: str) -> dict[str, list[str]]:
sections = {header: [] for header in SFT_HEADERS}
current: str | None = None
for raw_line in text.splitlines():
line = raw_line.strip()
lowered = line.lower()
if lowered in sections:
current = lowered
continue
if line.startswith("- "):
if current is None:
continue
sections[current].append(line[2:].strip())
return sections
def parse_answer_block(text: str) -> dict | None:
match = ANSWER_BLOCK_RE.search(text)
if match is None:
return None
try:
parsed = json.loads(match.group(1))
except json.JSONDecodeError:
return None
return parsed if isinstance(parsed, dict) else None
def parse_answer_text_block(text: str) -> str | None:
match = ANSWER_TEXT_BLOCK_RE.search(text)
if match is None:
return None
lines = [line.strip() for line in match.group(1).splitlines() if line.strip()]
if len(lines) != 1:
return None
return lines[0]
def validate_expected_answer(answer: object, row_id: str, path: str) -> list[str]:
errors: list[str] = []
if not isinstance(answer, dict):
return [f"{row_id}: {path} must be an object"]
if set(answer) != {"dominant", "irreducible", "reducible"}:
errors.append(f"{row_id}: {path} must contain exactly dominant, irreducible, reducible")
return errors
dominant = answer.get("dominant")
if not isinstance(dominant, str) or not dominant:
errors.append(f"{row_id}: {path}.dominant must be a non-empty string")
elif dominant not in PROCESS_IDS:
errors.append(f"{row_id}: unknown dominant process id {dominant!r}")
for key in ["irreducible", "reducible"]:
values = answer.get(key)
if not isinstance(values, list):
errors.append(f"{row_id}: {path}.{key} must be a list")
continue
for value in values:
if not isinstance(value, str):
errors.append(f"{row_id}: {path}.{key} contains a non-string id")
elif value not in PROCESS_IDS:
errors.append(f"{row_id}: unknown {key} process id {value!r}")
return errors
def validate_common(row: dict, row_id: str) -> list[str]:
errors: list[str] = []
required = {
"id",
"source_id",
"split",
"task_type",
"title",
"year",
"physics_target",
"broad_physics_area",
"final_state",
"process_catalog_version",
"expected_answer",
"metadata",
"target_type",
"target_category",
"target_process_id",
"target_background",
}
missing = required - set(row)
if missing:
errors.append(f"{row_id}: missing fields {sorted(missing)}")
if row.get("task_type") != TASK_TYPE:
errors.append(f"{row_id}: invalid task_type {row.get('task_type')!r}")
if row.get("split") not in SPLITS:
errors.append(f"{row_id}: invalid split {row.get('split')!r}")
if not isinstance(row.get("final_state"), dict):
errors.append(f"{row_id}: final_state must be an object")
if row.get("process_catalog_version") != PROCESS_CATALOG_VERSION:
errors.append(f"{row_id}: process_catalog_version must be {PROCESS_CATALOG_VERSION!r}")
errors.extend(validate_expected_answer(row.get("expected_answer"), row_id, "expected_answer"))
target_type = row.get("target_type")
target_category = row.get("target_category")
target_process_id = row.get("target_process_id")
target_background = row.get("target_background")
if target_type not in SFT_TARGET_TYPES:
errors.append(f"{row_id}: invalid target_type {target_type!r}")
elif target_category != SFT_TARGET_TYPES[target_type]:
errors.append(f"{row_id}: target_category {target_category!r} does not match target_type {target_type!r}")
if not isinstance(target_process_id, str) or target_process_id not in PROCESS_IDS:
errors.append(f"{row_id}: target_process_id must be a known process id")
if not isinstance(target_background, str) or not target_background:
errors.append(f"{row_id}: target_background must be a non-empty string")
row_expected = row.get("expected_answer")
if (
isinstance(row_expected, dict)
and isinstance(target_category, str)
and isinstance(target_process_id, str)
and target_process_id not in row_expected.get(target_category, [])
):
errors.append(f"{row_id}: target_process_id must appear in expected_answer.{target_category}")
metadata = row.get("metadata")
if not isinstance(metadata, dict):
errors.append(f"{row_id}: metadata must be an object")
return errors
if metadata.get("process_catalog_version") != PROCESS_CATALOG_VERSION:
errors.append(f"{row_id}: metadata.process_catalog_version must be {PROCESS_CATALOG_VERSION!r}")
if metadata.get("expected_answer") != row.get("expected_answer"):
errors.append(f"{row_id}: metadata.expected_answer must match row expected_answer")
for key in ["target_type", "target_category", "target_process_id", "target_background"]:
if metadata.get(key) != row.get(key):
errors.append(f"{row_id}: metadata.{key} must match row {key}")
for key in ["dominant_backgrounds", "irreducible_backgrounds", "reducible_backgrounds", "ranked_processes"]:
if not isinstance(metadata.get(key), list) or not metadata.get(key):
errors.append(f"{row_id}: metadata.{key} must be a non-empty list")
if isinstance(metadata.get("dominant_backgrounds"), list) and len(metadata["dominant_backgrounds"]) != 1:
errors.append(f"{row_id}: metadata.dominant_backgrounds must contain exactly one process")
row_dominant = row_expected.get("dominant") if isinstance(row_expected, dict) else None
if metadata.get("dominant_process_id") != row_dominant:
errors.append(f"{row_id}: metadata.dominant_process_id must match expected_answer.dominant")
if isinstance(target_category, str) and isinstance(target_background, str):
category_labels = metadata.get(f"{target_category}_backgrounds")
if isinstance(category_labels, list) and target_background not in category_labels:
errors.append(f"{row_id}: target_background must appear in metadata.{target_category}_backgrounds")
for key in ["irreducible_process_ids", "reducible_process_ids", "ranked_process_ids"]:
values = metadata.get(key)
if not isinstance(values, list):
errors.append(f"{row_id}: metadata.{key} must be a list")
continue
unknown = [value for value in values if value not in PROCESS_IDS]
if unknown:
errors.append(f"{row_id}: metadata.{key} contains unknown ids {unknown[:5]}")
return errors
def validate_messages(row: dict, row_id: str) -> list[str]:
errors: list[str] = []
messages = row.get("messages")
if not isinstance(messages, list) or len(messages) != 3:
return [f"{row_id}: messages must have exactly 3 entries"]
roles = [message.get("role") for message in messages if isinstance(message, dict)]
if roles != ["system", "user", "assistant"]:
errors.append(f"{row_id}: wrong message roles {roles}")
for idx, label in [(1, "user message"), (2, "assistant message")]:
content = str(messages[idx].get("content", "")) if isinstance(messages[idx], dict) else ""
match = TRAINABLE_EVIDENCE_TRACE_RE.search(content)
if match:
errors.append(f"{row_id}: {label} contains evidence trace text {match.group(0)!r}")
return errors
def validate_sft(row: dict) -> list[str]:
row_id = str(row.get("id", ""))
errors = validate_common(row, row_id)
errors.extend(validate_messages(row, row_id))
if errors:
return errors
answer = str(row["messages"][2]["content"])
lowered = answer.lower()
answer_matches = ANSWER_TEXT_BLOCK_RE.findall(answer)
if THINK_BLOCK_RE.search(answer):
errors.append(f"{row_id}: SFT answer must not contain a block")
if len(answer_matches) != 1:
errors.append(f"{row_id}: SFT answer must contain exactly one ... block")
if "" in lowered:
errors.append(f"{row_id}: SFT answer must not contain think tags")
if len(answer.split()) > 40:
errors.append(f"{row_id}: SFT answer is too long")
answer_text = parse_answer_text_block(answer)
if answer_text is None:
errors.append(f"{row_id}: SFT answer block must contain exactly one non-empty line")
elif answer_text.startswith("- "):
errors.append(f"{row_id}: SFT answer block must not use bullets")
else:
expected_label = row.get("target_background")
if answer_text != expected_label:
errors.append(f"{row_id}: SFT answer {answer_text!r} must match target background {expected_label!r}")
return errors
def validate_rl(row: dict) -> list[str]:
row_id = str(row.get("id", ""))
errors = validate_common(row, row_id)
required = {"prompt", "chosen_answer", "rejected_answer", "quality_note"}
missing = required - set(row)
if missing:
errors.append(f"{row_id}: missing RL fields {sorted(missing)}")
return errors
for field in ["prompt", "chosen_answer", "rejected_answer"]:
match = TRAINABLE_EVIDENCE_TRACE_RE.search(str(row.get(field, "")))
if match:
errors.append(f"{row_id}: {field} contains evidence trace text {match.group(0)!r}")
chosen = str(row.get("chosen_answer", ""))
lowered = chosen.lower()
if "" not in lowered or "" not in lowered:
errors.append(f"{row_id}: RL chosen_answer must contain a ... block")
else:
think_match = THINK_BLOCK_RE.search(chosen)
think_text = think_match.group(1).lower() if think_match else ""
for required_phrase in [
"same reconstructed final-state particles",
"irreducible backgrounds",
"fakes",
"reducible backgrounds",
]:
if required_phrase not in think_text:
errors.append(f"{row_id}: RL think block must mention {required_phrase!r}")
answer = parse_answer_block(chosen)
if answer is None:
errors.append(f"{row_id}: RL chosen_answer must contain parseable JSON")
else:
errors.extend(validate_expected_answer(answer, row_id, "chosen_answer answer"))
if answer != row.get("expected_answer"):
errors.append(f"{row_id}: chosen_answer JSON must exactly match expected_answer")
if str(row.get("chosen_answer", "")).strip() == str(row.get("rejected_answer", "")).strip():
errors.append(f"{row_id}: chosen_answer and rejected_answer are identical")
return errors
def main() -> int:
errors: list[str] = []
sft = read_jsonl(PROCESSED / "sft.jsonl")
rl = read_jsonl(PROCESSED / "rl.jsonl")
if not sft:
errors.append("sft.jsonl is empty")
if not rl:
errors.append("rl.jsonl is empty")
for name, rows in [("sft", sft), ("rl", rl)]:
ids = [row.get("id") for row in rows if row.get("id")]
dupes = [item for item, count in Counter(ids).items() if count > 1]
if dupes:
errors.append(f"{name}: duplicate ids: {dupes[:10]}")
sft_by_id = {row.get("id"): row for row in sft}
rl_by_id = {row.get("id"): row for row in rl}
if set(sft_by_id) != set(rl_by_id):
errors.append("SFT and RL ids do not match")
for row in sft:
errors.extend(validate_sft(row))
for row in rl:
errors.extend(validate_rl(row))
split_counts = Counter(row.get("split") for row in sft)
if "train" not in split_counts:
errors.append("No train examples found")
if not ({"val", "test"} & set(split_counts)):
errors.append("No validation/test examples found")
if errors:
print("Validation failed:")
for error in errors:
print(f"- {error}")
return 1
print("Validation passed")
print(f"SFT examples: {len(sft)}")
print(f"RL examples: {len(rl)}")
print(f"SFT split counts: {dict(split_counts)}")
return 0
if __name__ == "__main__":
raise SystemExit(main())