Download sdg/pipeline.py from fzzhang/svd-code: direct link, hf CLI and curl.
- Browser
- Download file 29.1 kB
-
https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/pipeline.py
- Command line
-
hf download hf://fzzhang/svd-code/sdg/pipeline.py
-
curl -L -o pipeline.py https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/pipeline.py
29.1 kB
| """ | |
| Main SDG pipeline orchestration. | |
| Stages: | |
| 1. Load HF dataset β extract instruction_seed | |
| 2. Init VLLMEngine | |
| 3. Generate N candidates per prompt, static-check, collect valid samples | |
| 4. Iterative UQ (max_rounds): cycle β factual β correctness, early-exit | |
| 5. Write output JSONL + stats.json + config.yaml | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import math | |
| import sys | |
| from dataclasses import dataclass, field, asdict | |
| from pathlib import Path | |
| from datasets import load_dataset | |
| from tqdm import tqdm | |
| from sdg.config import SDGConfig | |
| from sdg.inference import VLLMEngine | |
| from sdg.prompts import REASONING_INSTRUCTION | |
| from sdg.validation import ( | |
| collect_valid_samples, | |
| validate_batch_correctness, | |
| validate_batch_cycle, | |
| validate_batch_factual, | |
| validate_batch_simple, | |
| ) | |
| # Valid values for SDGConfig.selection_policy | |
| VALID_SELECTION_POLICIES = {"first_valid", "all_valid"} | |
| # ============================================================================= | |
| # TeeStream β duplicate stdout to a log file | |
| # ============================================================================= | |
| class TeeStream: | |
| """Write to both the original stream and a log file.""" | |
| def __init__(self, log_path: Path, original_stream, share_file=None): | |
| self.log_file = share_file if share_file is not None else open(log_path, "w") | |
| self._owns_file = share_file is None | |
| self.original = original_stream | |
| def write(self, data): | |
| self.original.write(data) | |
| self.log_file.write(data) | |
| self.log_file.flush() | |
| def flush(self): | |
| self.original.flush() | |
| self.log_file.flush() | |
| def fileno(self): | |
| return self.original.fileno() | |
| def isatty(self): | |
| return self.original.isatty() | |
| def close(self): | |
| if self._owns_file: | |
| self.log_file.close() | |
| # ============================================================================= | |
| # Row state tracker | |
| # ============================================================================= | |
| class RowState: | |
| row_id: int | |
| instruction_seed: str | |
| valid_samples: list[str] = field(default_factory=list) | |
| valid_indices: list[int] = field(default_factory=list) | |
| passed_sample: str | None = None # first_valid: single sample that passed all checks | |
| passed_samples: list[str] = field(default_factory=list) # all_valid: all samples that passed | |
| passed_round: int | None = None | |
| best_stage: int = 0 # deepest check reached: 0=none, 1=cycle, 2=factual, 3=correctness | |
| # ============================================================================= | |
| # Dataset loading | |
| # ============================================================================= | |
| def load_seeds(config: SDGConfig) -> list[str]: | |
| """Load instruction seeds from HuggingFace dataset.""" | |
| print(f"Loading dataset: {config.dataset_hf_id} (config={config.dataset_hf_config}, split={config.dataset_split}, revision={config.dataset_hf_revision})") | |
| ds = load_dataset(config.dataset_hf_id, name=config.dataset_hf_config, split=config.dataset_split, revision=config.dataset_hf_revision) | |
| if config.limit: | |
| ds = ds.select(range(min(config.limit, len(ds)))) | |
| print(f" Limited to {len(ds)} rows") | |
| seeds: list[str] = [] | |
| for row in ds: | |
| if config.dataset_format == "openthoughts4": | |
| convos = row.get("conversations", []) | |
| if convos: | |
| seeds.append(convos[0]["value"]) | |
| elif config.dataset_format == "messages": | |
| msgs = row.get("messages", []) | |
| if msgs: | |
| seeds.append(msgs[0]["content"]) | |
| else: | |
| raise ValueError(f"Unknown dataset_format: {config.dataset_format}") | |
| print(f" Extracted {len(seeds)} instruction seeds") | |
| return seeds | |
| # ============================================================================= | |
| # Stage 3: Generation + static check | |
| # ============================================================================= | |
| def generate_candidates( | |
| engine: VLLMEngine, | |
| seeds: list[str], | |
| config: SDGConfig, | |
| ) -> list[RowState]: | |
| """Generate N candidates per seed, run static check, build RowState list.""" | |
| rows: list[RowState] = [] | |
| total = len(seeds) | |
| processed = 0 | |
| for batch_start in range(0, total, config.gen_batch_size): | |
| batch_seeds = seeds[batch_start : batch_start + config.gen_batch_size] | |
| if config.check_boxed: | |
| prompts = [f"{seed}\n{REASONING_INSTRUCTION}\n" for seed in batch_seeds] | |
| else: | |
| prompts = [f"{seed}\n" for seed in batch_seeds] | |
| # generate N samples per prompt | |
| all_samples = engine.generate_multi(prompts) | |
| for i, (seed, samples) in enumerate(zip(batch_seeds, all_samples)): | |
| valid, indices = collect_valid_samples( | |
| samples, config.is_instruction_tuned, config.check_boxed | |
| ) | |
| if valid: | |
| rows.append( | |
| RowState( | |
| row_id=batch_start + i, | |
| instruction_seed=seed, | |
| valid_samples=valid, | |
| valid_indices=indices, | |
| ) | |
| ) | |
| processed += len(batch_seeds) | |
| print(f"[Generation] {processed}/{total} examples ({processed/total*100:.1f}%) | {len(rows)} valid") | |
| print( | |
| f" {len(rows)} / {len(seeds)} rows have at least one valid sample" | |
| ) | |
| return rows | |
| # ============================================================================= | |
| # Stage 4: Iterative UQ | |
| # ============================================================================= | |
| class RoundStats: | |
| round_idx: int = 0 | |
| pending: int = 0 | |
| exhausted: int = 0 | |
| cycle_passed: int = 0 | |
| cycle_failed: int = 0 | |
| factual_passed: int = 0 | |
| factual_failed: int = 0 | |
| correctness_passed: int = 0 | |
| correctness_failed: int = 0 | |
| newly_passed: int = 0 | |
| def _validate_round( | |
| engine: VLLMEngine, | |
| round_items: list[tuple[RowState, str]], | |
| config: SDGConfig, | |
| stats: RoundStats, | |
| ) -> tuple[list[tuple[RowState, str]], list[RowState], list[RowState], list[RowState]]: | |
| """Run cycle β factual β correctness on a batch of (row, sample) pairs. | |
| Returns: | |
| (correctness_passed_items, cycle_failed_rows, factual_failed_rows, correctness_failed_rows) | |
| """ | |
| # ββ cycle consistency ββββββββββββββββββββββββββββββββββββββββ | |
| cycle_items = [ | |
| (row.row_id, row.instruction_seed, sample) | |
| for row, sample in round_items | |
| ] | |
| cycle_results = validate_batch_cycle( | |
| engine, cycle_items, config.val_batch_size | |
| ) | |
| cycle_passed_items: list[tuple[RowState, str]] = [] | |
| cycle_failed_rows: list[RowState] = [] | |
| for row, sample in round_items: | |
| if cycle_results.get(row.row_id, False): | |
| cycle_passed_items.append((row, sample)) | |
| stats.cycle_passed += 1 | |
| row.best_stage = max(row.best_stage, 1) | |
| else: | |
| cycle_failed_rows.append(row) | |
| stats.cycle_failed += 1 | |
| print( | |
| f" Cycle: {stats.cycle_passed} passed, {stats.cycle_failed} failed" | |
| ) | |
| # ββ factual check (cycle-passed only) ββββββββββββββββββββββββ | |
| factual_passed_items: list[tuple[RowState, str]] = [] | |
| factual_failed_rows: list[RowState] = [] | |
| if cycle_passed_items: | |
| factual_items = [ | |
| (row.row_id, row.instruction_seed, sample) | |
| for row, sample in cycle_passed_items | |
| ] | |
| factual_results = validate_batch_factual( | |
| engine, factual_items, config.val_batch_size | |
| ) | |
| for row, sample in cycle_passed_items: | |
| if factual_results.get(row.row_id, False): | |
| factual_passed_items.append((row, sample)) | |
| stats.factual_passed += 1 | |
| row.best_stage = max(row.best_stage, 2) | |
| else: | |
| factual_failed_rows.append(row) | |
| stats.factual_failed += 1 | |
| print( | |
| f" Factual: {stats.factual_passed} passed, {stats.factual_failed} failed" | |
| ) | |
| # ββ correctness check (factual-passed only) βββββββββββββββββ | |
| correctness_passed_items: list[tuple[RowState, str]] = [] | |
| correctness_failed_rows: list[RowState] = [] | |
| if factual_passed_items: | |
| correctness_items = [ | |
| (row.row_id, row.instruction_seed, sample) | |
| for row, sample in factual_passed_items | |
| ] | |
| correctness_results = validate_batch_correctness( | |
| engine, correctness_items, config.val_batch_size | |
| ) | |
| for row, sample in factual_passed_items: | |
| if correctness_results.get(row.row_id, False): | |
| correctness_passed_items.append((row, sample)) | |
| stats.correctness_passed += 1 | |
| row.best_stage = max(row.best_stage, 3) | |
| else: | |
| correctness_failed_rows.append(row) | |
| stats.correctness_failed += 1 | |
| print( | |
| f" Correctness: {stats.correctness_passed} passed, " | |
| f"{stats.correctness_failed} failed" | |
| ) | |
| return correctness_passed_items, cycle_failed_rows, factual_failed_rows, correctness_failed_rows | |
| def run_uq_rounds( | |
| engine: VLLMEngine, | |
| rows: list[RowState], | |
| config: SDGConfig, | |
| ) -> tuple[list[RowState], list[RowState], list[RoundStats]]: | |
| """ | |
| Iterative UQ loop (first_valid policy). | |
| Tests one sample per row per round. Once a row passes all three checks | |
| it is moved to passed_rows and not tested again. | |
| Returns: | |
| (passed_rows, remaining_rows, per_round_stats) | |
| """ | |
| passed_rows: list[RowState] = [] | |
| pending_rows: list[RowState] = rows[:] | |
| exhausted_rows: list[RowState] = [] | |
| all_stats: list[RoundStats] = [] | |
| for round_idx in range(config.num_generations): | |
| if not pending_rows: | |
| print(f"Round {round_idx}: no pending rows β done early") | |
| break | |
| stats = RoundStats(round_idx=round_idx, pending=len(pending_rows)) | |
| print(f"\n{'='*60}") | |
| print(f"Round {round_idx}: {len(pending_rows)} pending rows") | |
| print(f"{'='*60}") | |
| # ββ extract sample at round_idx ββββββββββββββββββββββββββββββ | |
| round_items: list[tuple[RowState, str]] = [] # (row, current_sample) | |
| skipped: list[RowState] = [] | |
| for row in pending_rows: | |
| if round_idx < len(row.valid_samples): | |
| round_items.append((row, row.valid_samples[round_idx])) | |
| else: | |
| skipped.append(row) # exhausted samples | |
| if not round_items: | |
| print(f" No rows have a sample at index {round_idx} β stopping") | |
| break | |
| stats.exhausted = len(skipped) | |
| exhausted_rows.extend(skipped) | |
| print(f" {len(round_items)} rows have sample at index {round_idx}") | |
| if skipped: | |
| print(f" {len(skipped)} rows exhausted all samples (dropped)") | |
| # ββ validate βββββββββββββββββββββββββββββββββββββββββββββββββ | |
| correctness_passed, cycle_failed, factual_failed, correctness_failed = ( | |
| _validate_round(engine, round_items, config, stats) | |
| ) | |
| newly_passed: list[RowState] = [] | |
| for row, sample in correctness_passed: | |
| row.passed_sample = sample | |
| row.passed_round = round_idx | |
| newly_passed.append(row) | |
| stats.newly_passed = len(newly_passed) | |
| passed_rows.extend(newly_passed) | |
| # ββ remaining pending = all failures from this round βββββββββ | |
| pending_rows = cycle_failed + factual_failed + correctness_failed | |
| print( | |
| f" Round {round_idx} summary: {stats.newly_passed} newly passed, " | |
| f"{len(pending_rows)} still pending" | |
| ) | |
| all_stats.append(stats) | |
| # remaining = still pending after last round + exhausted across all rounds | |
| all_failed = pending_rows + exhausted_rows | |
| return passed_rows, all_failed, all_stats | |
| def _validate_round_simple( | |
| engine: VLLMEngine, | |
| round_items: list[tuple[RowState, str]], | |
| config: SDGConfig, | |
| stats: RoundStats, | |
| ) -> tuple[list[tuple[RowState, str]], list[RowState]]: | |
| """Single-prompt yes/no validation. Reuses correctness_* stat fields.""" | |
| items = [ | |
| (row.row_id, row.instruction_seed, sample) | |
| for row, sample in round_items | |
| ] | |
| results = validate_batch_simple(engine, items, config.val_batch_size) | |
| passed_items: list[tuple[RowState, str]] = [] | |
| failed_rows: list[RowState] = [] | |
| for row, sample in round_items: | |
| if results.get(row.row_id, False): | |
| passed_items.append((row, sample)) | |
| stats.correctness_passed += 1 | |
| row.best_stage = max(row.best_stage, 3) | |
| else: | |
| failed_rows.append(row) | |
| stats.correctness_failed += 1 | |
| # treat simple-fail as a correctness-stage failure for reporting | |
| row.best_stage = max(row.best_stage, 2) | |
| print( | |
| f" Simple: {stats.correctness_passed} passed, " | |
| f"{stats.correctness_failed} failed" | |
| ) | |
| return passed_items, failed_rows | |
| def run_simple_rounds( | |
| engine: VLLMEngine, | |
| rows: list[RowState], | |
| config: SDGConfig, | |
| ) -> tuple[list[RowState], list[RowState], list[RoundStats]]: | |
| """Iterative simple validation loop (first_valid policy).""" | |
| passed_rows: list[RowState] = [] | |
| pending_rows: list[RowState] = rows[:] | |
| exhausted_rows: list[RowState] = [] | |
| all_stats: list[RoundStats] = [] | |
| for round_idx in range(config.num_generations): | |
| if not pending_rows: | |
| print(f"Round {round_idx}: no pending rows β done early") | |
| break | |
| stats = RoundStats(round_idx=round_idx, pending=len(pending_rows)) | |
| print(f"\n{'='*60}") | |
| print(f"Round {round_idx}: {len(pending_rows)} pending rows") | |
| print(f"{'='*60}") | |
| round_items: list[tuple[RowState, str]] = [] | |
| skipped: list[RowState] = [] | |
| for row in pending_rows: | |
| if round_idx < len(row.valid_samples): | |
| round_items.append((row, row.valid_samples[round_idx])) | |
| else: | |
| skipped.append(row) | |
| if not round_items: | |
| print(f" No rows have a sample at index {round_idx} β stopping") | |
| break | |
| stats.exhausted = len(skipped) | |
| exhausted_rows.extend(skipped) | |
| print(f" {len(round_items)} rows have sample at index {round_idx}") | |
| if skipped: | |
| print(f" {len(skipped)} rows exhausted all samples (dropped)") | |
| passed_items, failed_rows = _validate_round_simple( | |
| engine, round_items, config, stats | |
| ) | |
| newly_passed: list[RowState] = [] | |
| for row, sample in passed_items: | |
| row.passed_sample = sample | |
| row.passed_round = round_idx | |
| newly_passed.append(row) | |
| stats.newly_passed = len(newly_passed) | |
| passed_rows.extend(newly_passed) | |
| pending_rows = failed_rows | |
| print( | |
| f" Round {round_idx} summary: {stats.newly_passed} newly passed, " | |
| f"{len(pending_rows)} still pending" | |
| ) | |
| all_stats.append(stats) | |
| all_failed = pending_rows + exhausted_rows | |
| return passed_rows, all_failed, all_stats | |
| def run_simple_rounds_all_valid( | |
| engine: VLLMEngine, | |
| rows: list[RowState], | |
| config: SDGConfig, | |
| ) -> tuple[list[RowState], list[RowState], list[RoundStats]]: | |
| """Iterative simple validation loop (all_valid policy).""" | |
| all_stats: list[RoundStats] = [] | |
| for round_idx in range(config.num_generations): | |
| round_items: list[tuple[RowState, str]] = [] | |
| for row in rows: | |
| if round_idx < len(row.valid_samples): | |
| round_items.append((row, row.valid_samples[round_idx])) | |
| if not round_items: | |
| print(f"Round {round_idx}: no rows have a sample at this index β stopping") | |
| break | |
| stats = RoundStats(round_idx=round_idx, pending=len(round_items)) | |
| print(f"\n{'='*60}") | |
| print(f"Round {round_idx}: {len(round_items)} rows") | |
| print(f"{'='*60}") | |
| passed_items, _ = _validate_round_simple( | |
| engine, round_items, config, stats | |
| ) | |
| for row, sample in passed_items: | |
| row.passed_samples.append(sample) | |
| if row.passed_round is None: | |
| row.passed_round = round_idx | |
| stats.newly_passed = len(passed_items) | |
| print( | |
| f" Round {round_idx} summary: {stats.newly_passed} samples passed" | |
| ) | |
| all_stats.append(stats) | |
| passed_rows = [row for row in rows if row.passed_samples] | |
| failed_rows = [row for row in rows if not row.passed_samples] | |
| return passed_rows, failed_rows, all_stats | |
| def run_uq_rounds_all_valid( | |
| engine: VLLMEngine, | |
| rows: list[RowState], | |
| config: SDGConfig, | |
| ) -> tuple[list[RowState], list[RowState], list[RoundStats]]: | |
| """ | |
| Iterative UQ loop (all_valid policy). | |
| Tests every valid sample for every row. A row is never removed early; | |
| all samples that pass all three checks are collected in row.passed_samples. | |
| Returns: | |
| (passed_rows, failed_rows, per_round_stats) | |
| """ | |
| all_stats: list[RoundStats] = [] | |
| for round_idx in range(config.num_generations): | |
| # ββ collect rows that have a sample at this index ββββββββββββ | |
| round_items: list[tuple[RowState, str]] = [] | |
| for row in rows: | |
| if round_idx < len(row.valid_samples): | |
| round_items.append((row, row.valid_samples[round_idx])) | |
| if not round_items: | |
| print(f"Round {round_idx}: no rows have a sample at this index β stopping") | |
| break | |
| stats = RoundStats(round_idx=round_idx, pending=len(round_items)) | |
| print(f"\n{'='*60}") | |
| print(f"Round {round_idx}: {len(round_items)} rows") | |
| print(f"{'='*60}") | |
| # ββ validate βββββββββββββββββββββββββββββββββββββββββββββββββ | |
| correctness_passed, _, _, _ = _validate_round( | |
| engine, round_items, config, stats, | |
| ) | |
| for row, sample in correctness_passed: | |
| row.passed_samples.append(sample) | |
| if row.passed_round is None: | |
| row.passed_round = round_idx | |
| stats.newly_passed = len(correctness_passed) | |
| print( | |
| f" Round {round_idx} summary: {stats.newly_passed} samples passed" | |
| ) | |
| all_stats.append(stats) | |
| passed_rows = [row for row in rows if row.passed_samples] | |
| failed_rows = [row for row in rows if not row.passed_samples] | |
| return passed_rows, failed_rows, all_stats | |
| # ============================================================================= | |
| # Stage 5: Write outputs | |
| # ============================================================================= | |
| def _build_stats( | |
| passed_rows: list[RowState], | |
| remaining_rows: list[RowState], | |
| round_stats: list[RoundStats], | |
| total_seeds: int, | |
| rows_after_static: int, | |
| config: SDGConfig | None = None, | |
| ) -> dict: | |
| """Build a well-organized stats dict for stats.json.""" | |
| static_failed = total_seeds - rows_after_static | |
| # Aggregate check totals across all rounds | |
| total_cycle_passed = sum(s.cycle_passed for s in round_stats) | |
| total_cycle_failed = sum(s.cycle_failed for s in round_stats) | |
| total_factual_passed = sum(s.factual_passed for s in round_stats) | |
| total_factual_failed = sum(s.factual_failed for s in round_stats) | |
| total_correctness_passed = sum(s.correctness_passed for s in round_stats) | |
| total_correctness_failed = sum(s.correctness_failed for s in round_stats) | |
| total_exhausted = sum(s.exhausted for s in round_stats) | |
| pass_rate = (len(passed_rows) / total_seeds * 100) if total_seeds else 0.0 | |
| total_passed_samples = sum(len(r.passed_samples) for r in passed_rows) | |
| # Example-level failure breakdown: categorize each failed example by the | |
| # deepest validation stage it ever reached across all rounds/samples. | |
| # best_stage 0 β never passed cycle (stuck at cycle) | |
| # best_stage 1 β passed cycle but never factual | |
| # best_stage 2 β passed factual but never correctness | |
| ex_failed_cycle = sum(1 for r in remaining_rows if r.best_stage == 0) | |
| ex_failed_factual = sum(1 for r in remaining_rows if r.best_stage == 1) | |
| ex_failed_correctness = sum(1 for r in remaining_rows if r.best_stage == 2) | |
| return { | |
| "summary": { | |
| "total_examples": total_seeds, | |
| "total_passed": len(passed_rows), | |
| "total_passed_samples": total_passed_samples, | |
| "total_failed": total_seeds - len(passed_rows), | |
| "pass_rate_pct": round(pass_rate, 2), | |
| }, | |
| "checks": { | |
| "static_check": { | |
| "passed": rows_after_static, | |
| "failed": static_failed, | |
| }, | |
| "cycle_consistency": { | |
| "passed": total_cycle_passed, | |
| "failed": total_cycle_failed, | |
| }, | |
| "factual_accuracy": { | |
| "passed": total_factual_passed, | |
| "failed": total_factual_failed, | |
| }, | |
| "correctness": { | |
| "passed": total_correctness_passed, | |
| "failed": total_correctness_failed, | |
| }, | |
| }, | |
| "failure_breakdown": { | |
| "failed_static_check": static_failed, | |
| "failed_cycle_consistency": total_cycle_failed, | |
| "failed_factual_accuracy": total_factual_failed, | |
| "failed_correctness": total_correctness_failed, | |
| "exhausted_samples": total_exhausted, | |
| "total_failed": total_seeds - len(passed_rows), | |
| }, | |
| "failure_breakdown_examples": { | |
| "failed_static_check": static_failed, | |
| "failed_cycle_consistency": ex_failed_cycle, | |
| "failed_factual_accuracy": ex_failed_factual, | |
| "failed_correctness": ex_failed_correctness, | |
| "total_failed": total_seeds - len(passed_rows), | |
| }, | |
| "rounds": [asdict(s) for s in round_stats], | |
| } | |
| def write_outputs( | |
| passed_rows: list[RowState], | |
| remaining_rows: list[RowState], | |
| round_stats: list[RoundStats], | |
| total_seeds: int, | |
| rows_after_static: int, | |
| config: SDGConfig, | |
| output_path: Path | None = None, | |
| ) -> None: | |
| out = output_path if output_path is not None else config.output_path | |
| out.mkdir(parents=True, exist_ok=True) | |
| # ββ output.jsonl βββββββββββββββββββββββββββββββββββββββββββββββββ | |
| jsonl_path = out / "output.jsonl" | |
| total_records = 0 | |
| with open(jsonl_path, "w") as f: | |
| for row in passed_rows: | |
| if config.check_boxed: | |
| user_content = row.instruction_seed.strip() + "\n" + REASONING_INSTRUCTION + "\n" | |
| else: | |
| user_content = row.instruction_seed.strip() + "\n" | |
| if config.selection_policy == "all_valid" and row.passed_samples: | |
| samples = row.passed_samples | |
| else: | |
| samples = [row.passed_sample] | |
| for sample in samples: | |
| record = { | |
| "messages": [ | |
| {"role": "user", "content": user_content}, | |
| {"role": "assistant", "content": sample}, | |
| ] | |
| } | |
| f.write(json.dumps(record, ensure_ascii=False) + "\n") | |
| total_records += 1 | |
| print(f"Wrote {total_records} records ({len(passed_rows)} examples) to {jsonl_path}") | |
| # ββ stats.json βββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| stats_path = out / "stats.json" | |
| stats_data = _build_stats( | |
| passed_rows, remaining_rows, round_stats, | |
| total_seeds, rows_after_static, config, | |
| ) | |
| with open(stats_path, "w") as f: | |
| json.dump(stats_data, f, indent=2) | |
| print(f"Wrote stats to {stats_path}") | |
| # ββ config.yaml copy βββββββββββββββββββββββββββββββββββββββββββββ | |
| config.to_yaml(out / "config.yaml") | |
| print(f"Wrote config to {out / 'config.yaml'}") | |
| # ============================================================================= | |
| # Public entry point | |
| # ============================================================================= | |
| def run_pipeline( | |
| config: SDGConfig, | |
| shard_id: int | None = None, | |
| num_shards: int | None = None, | |
| ) -> None: | |
| """Run the full SDG pipeline end-to-end.""" | |
| if config.selection_policy not in VALID_SELECTION_POLICIES: | |
| raise ValueError( | |
| f"Invalid selection_policy={config.selection_policy!r}, " | |
| f"must be one of {VALID_SELECTION_POLICIES}" | |
| ) | |
| is_sharded = shard_id is not None and num_shards is not None | |
| print(f"Run name: {config.run_name}") | |
| print(f"Output: {config.output_path}") | |
| if is_sharded: | |
| print(f"Shard: {shard_id}/{num_shards}") | |
| # ββ Per-shard logging via TeeStream ββββββββββββββββββββββββββββββ | |
| tee_out = None | |
| tee_err = None | |
| if is_sharded: | |
| log_dir = config.output_path / "logs" | |
| log_dir.mkdir(parents=True, exist_ok=True) | |
| log_path = log_dir / f"shard_{shard_id:03d}.log" | |
| tee_out = TeeStream(log_path, sys.stdout) | |
| tee_err = TeeStream(log_path, sys.stderr, share_file=tee_out.log_file) | |
| sys.stdout = tee_out | |
| sys.stderr = tee_err | |
| print(f"Logging to {log_path}") | |
| try: | |
| # Stage 1 β load data | |
| seeds = load_seeds(config) | |
| # ββ Shard slicing ββββββββββββββββββββββββββββββββββββββββββββ | |
| if is_sharded: | |
| total = len(seeds) | |
| chunk_size = math.ceil(total / num_shards) | |
| start = shard_id * chunk_size | |
| end = min(start + chunk_size, total) | |
| seeds = seeds[start:end] | |
| print(f"Shard {shard_id}: seeds [{start}:{end}] ({len(seeds)} examples)") | |
| # Stage 2 β init engine | |
| engine = VLLMEngine(config) | |
| # Stage 3 β generate candidates + static check | |
| rows = generate_candidates(engine, seeds, config) | |
| # Stage 4 β iterative UQ (skip if validation_type is None) | |
| rows_after_static = len(rows) | |
| if config.validation_type is None: | |
| print(f"\nSkipping validation (validation_type=null)") | |
| for row in rows: | |
| if config.selection_policy == "all_valid": | |
| row.passed_samples = list(row.valid_samples) | |
| else: | |
| row.passed_sample = row.valid_samples[0] | |
| row.passed_round = 0 | |
| row.best_stage = 3 | |
| passed = rows | |
| remaining = [] | |
| stats = [] | |
| elif config.validation_type == "simple": | |
| if config.selection_policy == "all_valid": | |
| passed, remaining, stats = run_simple_rounds_all_valid(engine, rows, config) | |
| else: | |
| passed, remaining, stats = run_simple_rounds(engine, rows, config) | |
| elif config.selection_policy == "all_valid": | |
| passed, remaining, stats = run_uq_rounds_all_valid(engine, rows, config) | |
| else: | |
| passed, remaining, stats = run_uq_rounds(engine, rows, config) | |
| # Stage 5 β write output | |
| total_seeds = len(seeds) | |
| output_path = None | |
| if is_sharded: | |
| output_path = config.output_path / "shards" / f"shard_{shard_id:03d}" | |
| write_outputs(passed, remaining, stats, total_seeds, rows_after_static, config, output_path=output_path) | |
| # Final summary | |
| print(f"\n{'='*60}") | |
| print("PIPELINE COMPLETE") | |
| print(f"{'='*60}") | |
| print(f"Total seeds: {total_seeds}") | |
| print(f"Total passed: {len(passed)}") | |
| print(f"Total dropped: {total_seeds - len(passed)}") | |
| if total_seeds: | |
| print(f"Pass rate: {len(passed) / total_seeds * 100:.1f}%") | |
| if output_path: | |
| print(f"Output dir: {output_path}") | |
| else: | |
| print(f"Output dir: {config.output_path}") | |
| finally: | |
| if tee_out is not None: | |
| sys.stdout = tee_out.original | |
| sys.stderr = tee_err.original | |
| tee_out.close() | |