Download sdg/validation.py from fzzhang/svd-code: direct link, hf CLI and curl.
- Browser
- Download file 6.91 kB
-
https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/validation.py
- Command line
-
hf download hf://fzzhang/svd-code/sdg/validation.py
-
curl -L -o validation.py https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/validation.py
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 | |