svd-code / sdg /pipeline.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
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
# =============================================================================
@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()