""" Validation logic: static checks, decision extraction, and batch UQ validators. Ported from experiments/self_instill/uq_verification.py and filter_with_uq.py. """ from __future__ import annotations import re from typing import TYPE_CHECKING from sdg.prompts import ( CYCLE_COMPARISON_PROMPT, CYCLE_QUESTION_GENERATION_PROMPT, FACTUAL_ERROR_PROMPT, SIMPLE_CORRECTNESS_PROMPT, TOTAL_CORRECTNESS_PROMPT, ) if TYPE_CHECKING: from sdg.inference import VLLMEngine # ============================================================================= # Regex patterns # ============================================================================= THINK_PATTERN = re.compile(r"\s*\S.*?\s*\S", re.DOTALL) BOXED_PATTERN = re.compile(r"\\boxed\{[^}]+\}") DECISION_RE = re.compile(r"\[\[\s*([YN])\s*\]\]", re.IGNORECASE) # ============================================================================= # Static check # ============================================================================= def passes_static_check(text: str, is_instruction_tuned: bool, check_boxed: bool = True) -> bool: if not text: return False if check_boxed and not BOXED_PATTERN.search(text): return False if is_instruction_tuned and not THINK_PATTERN.search(text): return False return True # ============================================================================= # Decision extraction # ============================================================================= def extract_decision_robust(text: str) -> bool: """Priority: [[Y]]/[[N]] → standalone Y/N → YES/NO → False.""" if not text: return False text_upper = text.upper() matches = list(DECISION_RE.finditer(text_upper)) if matches: return matches[-1].group(1) == "Y" standalone = re.findall(r"(? bool: if not vote_texts: return False return all(extract_decision_robust(t) for t in vote_texts) # ============================================================================= # Text helpers # ============================================================================= def extract_final_answer(answer: str) -> str: """Content after (or full text if no tag).""" if "" in answer: return answer.split("")[-1].strip() return answer def clean_inferred_question(text: str) -> str: """After , first line only.""" if not text: return "" if "" in text: text = text.split("")[-1].strip().split("\n")[0].strip() else: text = text.strip().split("\n")[0].strip() return text def collect_valid_samples( samples: list[str], is_instruction_tuned: bool, check_boxed: bool = True ) -> tuple[list[str], list[int]]: """Return (valid_samples, valid_indices) that pass static check.""" valid_samples: list[str] = [] valid_indices: list[int] = [] for idx, s in enumerate(samples): if passes_static_check(s, is_instruction_tuned, check_boxed): valid_samples.append(s) valid_indices.append(idx) return valid_samples, valid_indices # ============================================================================= # Batch validation functions # ============================================================================= def validate_batch_cycle( engine: "VLLMEngine", items: list[tuple[int, str, str]], val_batch_size: int, ) -> dict[int, bool]: """ Two-step cycle consistency: generate inferred Q, then compare. Args: items: list of (row_id, question, answer) Returns: {row_id: passed} """ results: dict[int, bool] = {} for batch_start in range(0, len(items), val_batch_size): batch = items[batch_start : batch_start + val_batch_size] # Step 1 — generate inferred questions gen_prompts = [ CYCLE_QUESTION_GENERATION_PROMPT.format(answer=extract_final_answer(a)) for _, _, a in batch ] inferred_raw = engine.generate_single(gen_prompts) inferred_clean = [clean_inferred_question(q) for q in inferred_raw] # Step 2 — compare compare_prompts = [ CYCLE_COMPARISON_PROMPT.format( original_question=q, inferred_question=iq ) for (_, q, _), iq in zip(batch, inferred_clean) ] vote_lists = engine.generate_with_votes(compare_prompts) for (row_id, _, _), votes in zip(batch, vote_lists): results[row_id] = check_unanimous(votes) return results def validate_batch_factual( engine: "VLLMEngine", items: list[tuple[int, str, str]], val_batch_size: int, ) -> dict[int, bool]: results: dict[int, bool] = {} for batch_start in range(0, len(items), val_batch_size): batch = items[batch_start : batch_start + val_batch_size] prompts = [ FACTUAL_ERROR_PROMPT.format(question=q, answer=extract_final_answer(a)) for _, q, a in batch ] vote_lists = engine.generate_with_votes(prompts) for (row_id, _, _), votes in zip(batch, vote_lists): results[row_id] = check_unanimous(votes) return results def validate_batch_correctness( engine: "VLLMEngine", items: list[tuple[int, str, str]], val_batch_size: int, ) -> dict[int, bool]: results: dict[int, bool] = {} for batch_start in range(0, len(items), val_batch_size): batch = items[batch_start : batch_start + val_batch_size] prompts = [ TOTAL_CORRECTNESS_PROMPT.format(question=q, answer=extract_final_answer(a)) for _, q, a in batch ] vote_lists = engine.generate_with_votes(prompts) for (row_id, _, _), votes in zip(batch, vote_lists): results[row_id] = check_unanimous(votes) return results def validate_batch_simple( engine: "VLLMEngine", items: list[tuple[int, str, str]], val_batch_size: int, ) -> dict[int, bool]: """Single yes/no judgement; passes only if all num_validation_votes votes say yes.""" results: dict[int, bool] = {} for batch_start in range(0, len(items), val_batch_size): batch = items[batch_start : batch_start + val_batch_size] prompts = [ SIMPLE_CORRECTNESS_PROMPT.format(question=q, answer=extract_final_answer(a)) for _, q, a in batch ] vote_lists = engine.generate_with_votes(prompts) for (row_id, _, _), votes in zip(batch, vote_lists): results[row_id] = check_unanimous(votes) return results