""" 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 # ============================================================================= @dataclass 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 # ============================================================================= @dataclass 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()