svd-code / sdg /validation.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
6.91 kB
"""
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"<think>\s*\S.*?</think>\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"(?<![A-Za-z])([YN])(?![A-Za-z])", text_upper)
if standalone:
return standalone[-1] == "Y"
yesno = re.findall(r"\b(YES|NO)\b", text_upper)
if yesno:
return yesno[-1] == "YES"
return False
def check_unanimous(vote_texts: list[str]) -> 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 </think> (or full text if no tag)."""
if "</think>" in answer:
return answer.split("</think>")[-1].strip()
return answer
def clean_inferred_question(text: str) -> str:
"""After </think>, first line only."""
if not text:
return ""
if "</think>" in text:
text = text.split("</think>")[-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