#!/usr/bin/env python3 """Generate and score immutable privileged-teacher gate records for ChartQA. Launch with ``torchrun`` to use one frozen 7B teacher replica per visible GPU. Each rank writes an isolated part; rank zero validates, orders, calibrates, and atomically publishes the final JSONL plus manifest. No training checkpoint is loaded or modified. """ from __future__ import annotations import argparse import json import math import os import sys import time from collections import defaultdict from pathlib import Path from typing import Any, Mapping, Sequence ROOT = Path(__file__).resolve().parents[1] if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) from data_utils.chart.evaluator import eval_teacher_probe_chart from data_utils.chart.teacher_gate import ( DEFAULT_CONFIDENCE_KEY, SCHEMA_VERSION, apply_gate_decision, calibrate_confidence_threshold, chartqa_dataset_fingerprint, chartqa_deplot_text, chartqa_image_key, chartqa_question, chartqa_reference_answer, chartqa_sample_fingerprint, contains_teacher_placeholder, default_manifest_path, record_hard_gate_passes, read_jsonl, read_manifest, summarize_teacher_gate_records, teacher_visible_chartqa_sample, validate_teacher_gate_cache, ) from data_utils.paths import ( local_pretrained_kwargs, resolve_image_path, resolve_model_path, validate_local_model_dir, ) from data_utils.rl_prompt import PROMPT_TEMPLATE from opsd_utils.privileged import build_privileged_context from opsd_utils.privileged.image_utils import load_rgb from opsd_utils.privileged.providers import split_teacher_response_prefix from opsd_utils.prompt_builder import _build_teacher_text def _parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--input", required=True, help="ChartQA JSON list") parser.add_argument("--output", required=True, help="Final teacher-gate JSONL") parser.add_argument("--split", required=True, choices=("train", "validation", "smoke")) parser.add_argument("--teacher-model", required=True, help="Complete local 7B checkpoint") parser.add_argument( "--processor-model", default="", help=( "Complete local processor/tokenizer directory. Defaults to --teacher-model; " "a separate path is useful when an offline weight snapshot omitted processor files" ), ) parser.add_argument( "--providers", default="visual_facts_deplot", help="Comma-separated privileged providers; format_only is forbidden", ) parser.add_argument("--batch-size", type=int, default=4) parser.add_argument("--max-new-tokens", type=int, default=96) parser.add_argument("--max-samples", type=int, default=0, help="0 means all rows") parser.add_argument( "--sample-strategy", choices=("head", "spread"), default="spread" ) parser.add_argument("--repetition-penalty", type=float, default=1.0) parser.add_argument( "--num-candidates", type=int, default=1, help="Gold-blind candidates per prompt; values above one enable sampling", ) parser.add_argument("--temperature", type=float, default=0.4) parser.add_argument("--top-p", type=float, default=0.95) parser.add_argument("--seed", type=int, default=42) parser.add_argument( "--source-selection", choices=("all", "hard_fail", "hard_pass"), default="all", help="Optionally pilot hard failures or hard passes from an existing full cache", ) parser.add_argument( "--selection-cache", default="", help="Full teacher-gate cache required by --source-selection=hard_fail", ) parser.add_argument( "--teacher-draft-cache", default="", help=( "Optional complete gold-blind teacher cache whose independently " "generated output is supplied as an untrusted self-review draft." ), ) parser.add_argument( "--merge-selection-cache", action="store_true", help=( "Publish a complete cache by replacing the selected rows in " "--selection-cache with their regenerated records. This is only " "valid for a complete hard_fail retry (no --max-samples)." ), ) parser.add_argument("--max-relative-change", type=float, default=0.05) parser.add_argument("--attn-implementation", choices=("sdpa", "flash_attention_2"), default="sdpa") parser.add_argument("--confidence-key", default=DEFAULT_CONFIDENCE_KEY) parser.add_argument("--calibrate", action="store_true") parser.add_argument("--target-precision", type=float, default=0.90) parser.add_argument("--min-coverage", type=float, default=0.10) parser.add_argument( "--calibration-manifest", default="", help="Manifest whose calibrated threshold is applied to this split", ) parser.add_argument( "--require-real-deplot", action=argparse.BooleanOptionalAction, default=True, ) parser.add_argument( "--human-only", action=argparse.BooleanOptionalAction, default=True, help=( "Filter human_or_machine!=0 rows (train default). Use --no-human-only " "for validation calibration so it matches the 1920-row evaluator split." ), ) parser.add_argument("--expected-samples", type=int, default=0) parser.add_argument("--log-every", type=int, default=10) return parser.parse_args(argv) def _distributed_context() -> tuple[Any, int, int, int]: import torch import torch.distributed as dist rank = int(os.environ.get("RANK", "0")) world_size = int(os.environ.get("WORLD_SIZE", "1")) local_rank = int(os.environ.get("LOCAL_RANK", "0")) if not torch.cuda.is_available(): raise RuntimeError("privileged teacher preprocessing requires CUDA") torch.cuda.set_device(local_rank) if world_size > 1 and not dist.is_initialized(): dist.init_process_group(backend="nccl") return torch, rank, world_size, local_rank def _barrier(world_size: int) -> None: if world_size <= 1: return import torch.distributed as dist dist.barrier() def _load_rows(path: str, *, human_only: bool) -> list[dict[str, Any]]: with Path(path).open(encoding="utf-8") as handle: payload = json.load(handle) if isinstance(payload, dict): payload = payload.get("data", payload.get("examples", [])) if not isinstance(payload, list) or not payload: raise ValueError("teacher gate input must be a non-empty JSON list") rows = [dict(row) for row in payload if isinstance(row, Mapping)] if human_only: rows = [row for row in rows if row.get("human_or_machine", 0) == 0] if not rows: raise ValueError("teacher gate input has no effective ChartQA rows") return rows def _selected_indices(total: int, maximum: int, strategy: str) -> list[int]: if maximum <= 0 or maximum >= total: return list(range(total)) if strategy == "head" or maximum == 1: return list(range(maximum)) # Deterministic coverage of the complete source ordering, including both # endpoints. Duplicates from rounding are filled by the first unused rows. selected = { int(round(position * (total - 1) / (maximum - 1))) for position in range(maximum) } if len(selected) < maximum: selected.update(index for index in range(total) if index not in selected) return sorted(selected)[:maximum] def _prepare_sample(raw: Mapping[str, Any]) -> dict[str, Any]: question = chartqa_question(raw) answer = chartqa_reference_answer(raw) image = raw.get("image") or raw.get("image_path") if not question or not answer or not image: raise ValueError("ChartQA teacher row is missing question, answer, or image") sample = dict(raw) sample["question"] = question sample["question_wo_prompt"] = question sample["answer"] = answer sample["reference_answer"] = answer sample["image"] = resolve_image_path(str(image)) sample["prompt"] = PROMPT_TEMPLATE.format(question=question) return sample def _build_job( sample: Mapping[str, Any], *, providers: list[str], source_index: int, teacher_draft: str = "", ) -> dict[str, Any]: if "format_only" in providers: raise ValueError("teacher gate preprocessing forbids the format_only provider") opsd_config = { "text_include_gold": False, "privileged_profile": "hybrid", "teacher_probe": {"enabled": False}, } teacher_sample = teacher_visible_chartqa_sample(sample) if str(teacher_draft or "").strip(): teacher_sample["_teacher_generated_draft"] = str(teacher_draft).strip() suffix, teacher_images = build_privileged_context( teacher_sample, providers, privileged_profile="hybrid", opsd_config=opsd_config, ) suffix, response_prefix = split_teacher_response_prefix(suffix) if response_prefix.strip(): raise ValueError("teacher gate prompt unexpectedly contains a response prefix") if not teacher_images: image = load_rgb(teacher_sample.get("image")) teacher_images = [image] if image is not None else [] prompt = _build_teacher_text(str(sample["prompt"]), suffix) if contains_teacher_placeholder(prompt): raise ValueError(f"teacher gate prompt {source_index} contains a forbidden placeholder") return { "source_index": int(source_index), "sample": dict(sample), "prompt": prompt, "images": teacher_images, "deplot_present": bool(chartqa_deplot_text(sample)), "teacher_draft": str(teacher_draft or "").strip(), } def _load_teacher(args: argparse.Namespace, torch: Any, local_rank: int) -> tuple[Any, Any]: from transformers import AutoProcessor, LlavaOnevisionForConditionalGeneration model_path = validate_local_model_dir(resolve_model_path(args.teacher_model), role="teacher") if not os.path.isdir(model_path): raise ValueError( "teacher preprocessing is offline-only; --teacher-model must be a local directory" ) processor_source = args.processor_model or args.teacher_model processor_path = resolve_model_path(processor_source) if not os.path.isdir(processor_path): raise ValueError( "teacher preprocessing is offline-only; --processor-model must be a local directory" ) local_kw = local_pretrained_kwargs(model_path) processor = AutoProcessor.from_pretrained( processor_path, **local_pretrained_kwargs(processor_path), ) processor.tokenizer.padding_side = "left" model = LlavaOnevisionForConditionalGeneration.from_pretrained( model_path, torch_dtype=torch.bfloat16, low_cpu_mem_usage=True, attn_implementation=args.attn_implementation, **local_kw, ).to(torch.device("cuda", local_rank)) model.eval() model.requires_grad_(False) return model, processor def _batch_signature(processed: Mapping[str, Any]) -> tuple[Any, ...]: pixel_values = processed.get("pixel_values") image_sizes = processed.get("image_sizes") return ( tuple(pixel_values.shape) if pixel_values is not None else (), tuple(image_sizes.shape) if image_sizes is not None else (), int(processed.get("batch_num_images", 1) or 1), ) def _valid_generated_ids( ids: Sequence[int], *, eos_token_id: int | None, pad_token_id: int | None ) -> tuple[list[int], bool]: valid: list[int] = [] ended = False for token in ids: token = int(token) if eos_token_id is not None and token == int(eos_token_id): ended = True break if pad_token_id is not None and token == int(pad_token_id): break valid.append(token) return valid, ended def _find_subsequence(haystack: Sequence[int], needle: Sequence[int]) -> list[int]: if not needle or len(needle) > len(haystack): return [] for start in range(len(haystack) - len(needle), -1, -1): if list(haystack[start : start + len(needle)]) == list(needle): return list(range(start, start + len(needle))) return [] def _answer_token_indices(tokenizer: Any, token_ids: list[int], answer: str) -> list[int]: answer = str(answer or "").strip() if not token_ids or not answer: return [] for variant in (answer, f" {answer}", f"\n{answer}"): encoded = tokenizer.encode(variant, add_special_tokens=False) found = _find_subsequence(token_ids, encoded) if found: return found # Character overlap fallback handles contextual BPE merges around the # colon/space preceding a short answer. decoded = tokenizer.decode(token_ids, skip_special_tokens=True) start = decoded.lower().rfind(answer.lower()) if start < 0: return list(range(len(token_ids))) if decoded.strip().lower() == answer.lower() else [] end = start + len(answer) prefix_lengths = [0] for stop in range(1, len(token_ids) + 1): prefix_lengths.append( len(tokenizer.decode(token_ids[:stop], skip_special_tokens=True)) ) overlap = [ index for index in range(len(token_ids)) if prefix_lengths[index] < end and prefix_lengths[index + 1] > start ] return overlap def _confidence_from_scores( torch: Any, scores: Sequence[Any], row: int, token_ids: Sequence[int], answer_indices: Sequence[int], ) -> dict[str, float]: log_probabilities: list[float] = [] probabilities: list[float] = [] entropies: list[float] = [] margins: list[float] = [] vocab_size = 0 for position in answer_indices: if position >= len(scores) or position >= len(token_ids): continue logits = scores[position][row].float() vocab_size = int(logits.numel()) log_probs = torch.log_softmax(logits, dim=-1) token_id = int(token_ids[position]) selected_log_prob = float(log_probs[token_id].item()) probs = torch.softmax(logits, dim=-1) selected_probability = float(probs[token_id].item()) entropy = float(torch.special.entr(probs).sum().item()) top_two = torch.topk(probs, k=min(2, vocab_size)).values margin = float( (top_two[0] - top_two[1]).item() if top_two.numel() > 1 else top_two[0].item() ) log_probabilities.append(selected_log_prob) probabilities.append(selected_probability) entropies.append(entropy) margins.append(margin) if not probabilities: return {} mean_log_probability = sum(log_probabilities) / len(log_probabilities) mean_entropy = sum(entropies) / len(entropies) return { "answer_token_count": float(len(probabilities)), "answer_token_geometric_mean_probability": float(math.exp(mean_log_probability)), "answer_token_mean_probability": float(sum(probabilities) / len(probabilities)), "answer_token_min_probability": float(min(probabilities)), "answer_token_mean_log_probability": float(mean_log_probability), "answer_token_mean_entropy": float(mean_entropy), "answer_token_mean_normalized_entropy": float( mean_entropy / math.log(max(vocab_size, 2)) ), "answer_token_mean_top1_margin": float(sum(margins) / len(margins)), } def _candidate_confidence(record: Mapping[str, Any]) -> float: confidence = record.get("confidence") return ( float(confidence.get(DEFAULT_CONFIDENCE_KEY, 0.0) or 0.0) if isinstance(confidence, Mapping) else 0.0 ) def _candidate_structurally_valid(record: Mapping[str, Any]) -> bool: return bool( record.get("parse_failed") is False and record.get("placeholder") is False and record.get("clipped") is False and record.get("answer_span_found") is True ) def _candidate_rank_key(record: Mapping[str, Any]) -> tuple[int, float, int, int]: """Gold-independent fallback preference for one generated candidate.""" return ( int(_candidate_structurally_valid(record)), _candidate_confidence(record), -int(record.get("generated_token_count", 0) or 0), -int(record.get("candidate_index", 0) or 0), ) def _select_gold_blind_candidate( candidates: Sequence[Mapping[str, Any]], ) -> dict[str, Any]: """Select one candidate by self-consistency without consulting correctness. Candidate correctness is attached by the evaluator before this function so it can be retained for audit, but it must never influence the selected index. The final gold gate is applied only to this independently selected answer. This prevents pass@N from becoming an oracle label selector. """ if not candidates: raise ValueError("cannot select from zero teacher candidates") grouped: dict[str, list[int]] = defaultdict(list) for index, candidate in enumerate(candidates): if not _candidate_structurally_valid(candidate): continue answer_key = " ".join( str(candidate.get("parsed_answer") or "").strip().casefold().split() ) if answer_key: grouped[answer_key].append(index) if grouped: winning_key = max( grouped, key=lambda key: ( len(grouped[key]), sum(_candidate_confidence(candidates[i]) for i in grouped[key]), max(_candidate_confidence(candidates[i]) for i in grouped[key]), any( str(candidates[i].get("candidate_mode") or "") == "greedy" for i in grouped[key] ), -min(grouped[key]), ), ) selected_index = max( grouped[winning_key], key=lambda i: ( _candidate_confidence(candidates[i]), str(candidates[i].get("candidate_mode") or "") == "greedy", -int(candidates[i].get("generated_token_count", 0) or 0), -i, ), ) else: selected_index = max( range(len(candidates)), key=lambda i: _candidate_rank_key( {**dict(candidates[i]), "candidate_index": i} ), ) selected = dict(candidates[selected_index]) selected["selection_policy"] = "gold_blind_plurality_confidence_v1" selected["candidate_count"] = len(candidates) selected["correct_candidate_count"] = sum( int(candidate.get("teacher_correct") is True) for candidate in candidates ) selected["hard_pass_candidate_count"] = sum( int(record_hard_gate_passes(candidate)) for candidate in candidates ) selected["selected_candidate_index"] = int(selected_index) selected["candidate_audit"] = [ { "candidate_index": int(index), "candidate_mode": str(candidate.get("candidate_mode") or "unknown"), "teacher_output": str(candidate.get("teacher_output") or ""), "parsed_answer": str(candidate.get("parsed_answer") or ""), "teacher_correct": bool(candidate.get("teacher_correct") is True), "parse_failed": bool(candidate.get("parse_failed") is True), "placeholder": bool(candidate.get("placeholder") is True), "clipped": bool(candidate.get("clipped") is True), "answer_span_found": bool(candidate.get("answer_span_found") is True), "confidence": dict(candidate.get("confidence") or {}), } for index, candidate in enumerate(candidates) ] return selected def _merge_retry_records( base_records: Sequence[Mapping[str, Any]], retry_records: Sequence[Mapping[str, Any]], *, expected_retry_indices: Sequence[int], ) -> list[dict[str, Any]]: """Replace a complete cache's retried rows without changing row order.""" if [int(record.get("source_index", -1)) for record in base_records] != list( range(len(base_records)) ): raise ValueError("selection cache records are not a complete ordered dataset") expected = [int(index) for index in expected_retry_indices] if len(expected) != len(set(expected)): raise ValueError("retry source indices contain duplicates") replacements = { int(record.get("source_index", -1)): dict(record) for record in retry_records } if sorted(replacements) != sorted(expected) or len(replacements) != len(retry_records): raise ValueError("retry records are incomplete or duplicated") if any(index < 0 or index >= len(base_records) for index in replacements): raise ValueError("retry source index is outside the selection cache") return [ replacements.get(index, dict(base_record)) for index, base_record in enumerate(base_records) ] def _generate_compatible_batch( *, torch: Any, model: Any, processor: Any, processed_batches: list[Mapping[str, Any]], jobs: list[Mapping[str, Any]], args: argparse.Namespace, ) -> list[dict[str, Any]]: from opsd_utils.teacher_batching import ( model_inference_device, move_batch_num_images_to_model_device, move_pixel_values_to_model_device, stack_teacher_processor_batches, ) from reward_utils.teacher_generate import _align_stacked_batch stacked = stack_teacher_processor_batches(processor, processed_batches) aligned_ids, aligned_mask, pixel_values, image_sizes, batch_num_images = _align_stacked_batch( model, processor, stacked ) if isinstance(pixel_values, list): if len(jobs) == 1: raise RuntimeError("single teacher row retained incompatible visual tensors") outputs: list[dict[str, Any]] = [] for processed, job in zip(processed_batches, jobs): outputs.extend( _generate_compatible_batch( torch=torch, model=model, processor=processor, processed_batches=[processed], jobs=[job], args=args, ) ) return outputs device = model_inference_device(model) aligned_ids = aligned_ids.to(device) aligned_mask = aligned_mask.to(device) pixel_values = move_pixel_values_to_model_device(model, pixel_values) batch_num_images = move_batch_num_images_to_model_device(model, batch_num_images) if hasattr(image_sizes, "to"): image_sizes = image_sizes.to(device) forward: dict[str, Any] = { "input_ids": aligned_ids, "attention_mask": aligned_mask, } if pixel_values is not None: forward.update( pixel_values=pixel_values, image_sizes=image_sizes, batch_num_images=batch_num_images, ) prompt_length = int(aligned_ids.shape[1]) num_candidates = int(args.num_candidates) generation_runs: list[tuple[str, Any, tuple[Any, ...], int]] = [] def generate_run(mode: str, count: int) -> None: generation_kwargs: dict[str, Any] = { "max_new_tokens": int(args.max_new_tokens), "do_sample": mode == "sampled", "num_return_sequences": int(count), "repetition_penalty": float(args.repetition_penalty), "pad_token_id": processor.tokenizer.pad_token_id, "eos_token_id": processor.tokenizer.eos_token_id, "return_dict_in_generate": True, "output_scores": True, } if mode == "sampled": generation_kwargs.update( temperature=float(args.temperature), top_p=float(args.top_p), ) with torch.inference_mode(): generated = model.generate(**forward, **generation_kwargs) scores = tuple(generated.scores) new_ids = generated.sequences[ :, prompt_length : prompt_length + len(scores) ] generation_runs.append((mode, new_ids, scores, int(count))) # Candidate zero is always deterministic. Full-cache regeneration can # therefore never lose a row merely because sampling omitted the teacher's # greedy answer; additional candidates only increase pass@N coverage. generate_run("greedy", 1) if num_candidates > 1: generate_run("sampled", num_candidates - 1) results: list[dict[str, Any]] = [] for row, job in enumerate(jobs): candidates: list[dict[str, Any]] = [] for mode, new_ids, scores, run_count in generation_runs: for run_candidate_index in range(run_count): generated_row = row * run_count + run_candidate_index valid_ids, ended_with_eos = _valid_generated_ids( new_ids[generated_row].detach().cpu().tolist(), eos_token_id=processor.tokenizer.eos_token_id, pad_token_id=processor.tokenizer.pad_token_id, ) text = processor.tokenizer.decode( valid_ids, skip_special_tokens=True ).strip() score, parsed = eval_teacher_probe_chart( text, chartqa_reference_answer(job["sample"]), float(args.max_relative_change), answer_flag="answer:", ) answer_indices = _answer_token_indices( processor.tokenizer, valid_ids, parsed.answer ) confidence = _confidence_from_scores( torch, scores, generated_row, valid_ids, answer_indices ) clipped = bool( not ended_with_eos and len(scores) >= int(args.max_new_tokens) ) candidates.append( { "schema_version": SCHEMA_VERSION, "source_index": int(job["source_index"]), "sample_fingerprint": chartqa_sample_fingerprint(job["sample"]), "question": chartqa_question(job["sample"]), "image": chartqa_image_key(job["sample"]), "reference_answer": chartqa_reference_answer(job["sample"]), "deplot_present": bool(job["deplot_present"]), "teacher_draft": str(job.get("teacher_draft") or ""), "candidate_mode": mode, "teacher_output": text, "parsed_answer": parsed.answer, "score": float(score), "teacher_correct": bool(score > 0.0), "parse_failed": bool(parsed.parse_failed), "has_answer_flag": bool(parsed.has_answer_flag), "placeholder": contains_teacher_placeholder(text), "clipped": clipped, "ended_with_eos": bool(ended_with_eos), "generated_token_count": len(valid_ids), "answer_span_found": bool(answer_indices), "answer_token_indices": list(answer_indices), "confidence": confidence, } ) results.append(_select_gold_blind_candidate(candidates)) generation_runs.clear() return results def _generate_jobs( *, torch: Any, model: Any, processor: Any, jobs: list[Mapping[str, Any]], args: argparse.Namespace, rank: int, ) -> list[dict[str, Any]]: from opsd_utils.teacher_batching import process_teacher_sample records: list[dict[str, Any]] = [] started = time.monotonic() for outer_start in range(0, len(jobs), max(1, int(args.batch_size))): outer_jobs = jobs[outer_start : outer_start + max(1, int(args.batch_size))] prepared: list[Mapping[str, Any]] = [] for job in outer_jobs: prepared.append( process_teacher_sample(processor, job["prompt"], job["images"]) ) groups: dict[tuple[Any, ...], list[int]] = defaultdict(list) for index, processed in enumerate(prepared): groups[_batch_signature(processed)].append(index) ordered_results: dict[int, dict[str, Any]] = {} for group_indices in groups.values(): group_records = _generate_compatible_batch( torch=torch, model=model, processor=processor, processed_batches=[prepared[index] for index in group_indices], jobs=[outer_jobs[index] for index in group_indices], args=args, ) for index, record in zip(group_indices, group_records): ordered_results[index] = record records.extend(ordered_results[index] for index in range(len(outer_jobs))) completed = min(outer_start + len(outer_jobs), len(jobs)) if completed == len(jobs) or completed % max(1, int(args.log_every)) == 0: elapsed = time.monotonic() - started rate = completed / elapsed if elapsed > 0 else 0.0 print( f"[teacher-gate][rank={rank}] {completed}/{len(jobs)} " f"rows elapsed={elapsed:.1f}s rate={rate:.3f}/s", flush=True, ) return records def _atomic_write_jsonl(path: Path, rows: Sequence[Mapping[str, Any]]) -> None: path.parent.mkdir(parents=True, exist_ok=True) temporary = path.with_name(f".{path.name}.tmp.{os.getpid()}") with temporary.open("w", encoding="utf-8") as handle: for row in rows: handle.write(json.dumps(dict(row), ensure_ascii=False, sort_keys=True) + "\n") os.replace(temporary, path) def _atomic_write_json(path: Path, payload: Mapping[str, Any]) -> None: path.parent.mkdir(parents=True, exist_ok=True) temporary = path.with_name(f".{path.name}.tmp.{os.getpid()}") with temporary.open("w", encoding="utf-8") as handle: json.dump(dict(payload), handle, ensure_ascii=False, sort_keys=True, indent=2) handle.write("\n") os.replace(temporary, path) def _calibration(args: argparse.Namespace, records: Sequence[Mapping[str, Any]]) -> dict[str, Any]: if args.calibration_manifest: manifest = read_manifest(args.calibration_manifest) calibration = manifest.get("calibration") if not isinstance(calibration, Mapping): raise ValueError("calibration manifest has no calibration block") if str(calibration.get("confidence_key")) != str(args.confidence_key): raise ValueError("calibration confidence key does not match this run") result = dict(calibration) result["source_manifest"] = str(Path(args.calibration_manifest).resolve()) return result if args.calibrate: return calibrate_confidence_threshold( records, confidence_key=str(args.confidence_key), target_precision=float(args.target_precision), min_coverage=float(args.min_coverage), ).to_dict() return { "confidence_key": str(args.confidence_key), "threshold": 0.0, "target_precision": None, "achieved_precision": None, "retained_count": len(records), "candidate_count": len(records), "total_count": len(records), "coverage": 1.0, "smoke_only": True, } def main(argv: Sequence[str] | None = None) -> int: args = _parse_args(argv) if ( args.batch_size <= 0 or args.max_new_tokens <= 0 or args.max_samples < 0 or args.num_candidates <= 0 ): raise ValueError("batch-size/max-new-tokens must be positive and max-samples non-negative") if args.num_candidates > 1 and not (0.0 < float(args.temperature)): raise ValueError("temperature must be positive when num-candidates > 1") if not 0.0 < float(args.top_p) <= 1.0: raise ValueError("top-p must be in (0, 1]") if abs(float(args.repetition_penalty) - 1.0) > 1e-12: raise ValueError("teacher gate greedy generation requires repetition_penalty=1.0") providers = [item.strip() for item in str(args.providers).split(",") if item.strip()] if not providers or "format_only" in providers: raise ValueError("providers must be non-empty and must not contain format_only") gold_blind_provider_allowlist = { "visual_facts_deplot", "chartqa_gold_blind", "chartqa_gold_blind_verify", "chartqa_self_review", "crop", } unsafe_providers = sorted(set(providers) - gold_blind_provider_allowlist) if unsafe_providers: raise ValueError( "teacher gate preprocessing accepts only gold-blind providers; " f"unsafe={unsafe_providers}" ) raw_rows = _load_rows(args.input, human_only=bool(args.human_only)) if args.expected_samples and len(raw_rows) != int(args.expected_samples): raise ValueError( f"effective input rows={len(raw_rows)}, expected={int(args.expected_samples)}" ) candidate_indices = list(range(len(raw_rows))) selection_records: list[dict[str, Any]] | None = None selection_manifest: dict[str, Any] | None = None if args.source_selection in {"hard_fail", "hard_pass"}: if not args.selection_cache: raise ValueError( "--source-selection hard_fail/hard_pass requires --selection-cache" ) selection_records, selection_manifest = validate_teacher_gate_cache( raw_rows, cache_path=args.selection_cache, require_validation_calibration=False, ) select_passes = args.source_selection == "hard_pass" candidate_indices = [ index for index, record in enumerate(selection_records) if record_hard_gate_passes(record) is select_passes ] if not candidate_indices: raise ValueError( f"selection cache contains no rows for {args.source_selection}" ) if args.merge_selection_cache: if args.source_selection != "hard_fail" or not args.selection_cache: raise ValueError( "--merge-selection-cache requires --source-selection hard_fail " "and --selection-cache" ) if int(args.max_samples) != 0: raise ValueError( "--merge-selection-cache requires a complete retry with --max-samples 0" ) if selection_records is None or selection_manifest is None: raise ValueError("selection cache was not loaded for merge") if list(selection_manifest.get("providers") or []) != providers: raise ValueError( "retry providers must exactly match the selection cache providers: " f"retry={providers}, base={selection_manifest.get('providers')}" ) draft_records: list[dict[str, Any]] | None = None draft_manifest: dict[str, Any] | None = None if args.teacher_draft_cache: if "chartqa_self_review" not in providers: raise ValueError( "--teacher-draft-cache requires the chartqa_self_review provider" ) draft_records, draft_manifest = validate_teacher_gate_cache( raw_rows, cache_path=args.teacher_draft_cache, require_validation_calibration=False, ) elif "chartqa_self_review" in providers: raise ValueError( "chartqa_self_review requires --teacher-draft-cache" ) torch, rank, world_size, local_rank = _distributed_context() selected_positions = _selected_indices( len(candidate_indices), int(args.max_samples), args.sample_strategy ) indices = [candidate_indices[position] for position in selected_positions] selected_samples = [_prepare_sample(raw_rows[index]) for index in indices] if args.require_real_deplot: missing = [ indices[position] for position, sample in enumerate(selected_samples) if not chartqa_deplot_text(sample) ] if missing: raise ValueError( f"selected teacher rows lack real DePlot evidence: count={len(missing)}, " f"examples={missing[:8]}" ) local_positions = list(range(rank, len(indices), world_size)) local_jobs = [ _build_job( selected_samples[position], providers=providers, source_index=indices[position], teacher_draft=( str(draft_records[indices[position]].get("teacher_output") or "") if draft_records is not None else "" ), ) for position in local_positions ] print( f"[teacher-gate][rank={rank}/{world_size}] device=cuda:{local_rank} " f"local_rows={len(local_jobs)} total_selected={len(indices)} providers={providers}", flush=True, ) model, processor = _load_teacher(args, torch, local_rank) torch.manual_seed(int(args.seed) + rank) torch.cuda.manual_seed_all(int(args.seed) + rank) records = _generate_jobs( torch=torch, model=model, processor=processor, jobs=local_jobs, args=args, rank=rank, ) output_path = Path(args.output).resolve() parts_dir = Path(f"{output_path}.parts") part_path = parts_dir / f"world{world_size:03d}-rank{rank:03d}.jsonl" _atomic_write_jsonl(part_path, records) _barrier(world_size) if rank == 0: merged: list[dict[str, Any]] = [] for part_rank in range(world_size): part = parts_dir / f"world{world_size:03d}-rank{part_rank:03d}.jsonl" if not part.is_file(): raise FileNotFoundError(f"missing teacher gate rank part: {part}") with part.open(encoding="utf-8") as handle: for line in handle: if line.strip(): merged.append(json.loads(line)) merged.sort(key=lambda record: int(record["source_index"])) if [int(record["source_index"]) for record in merged] != indices: raise ValueError("merged teacher gate source indices are incomplete or duplicated") retry_record_count = len(merged) recovered_rows = sum( int(record_hard_gate_passes(record)) for record in merged ) if args.merge_selection_cache: assert selection_records is not None merged = _merge_retry_records( selection_records, merged, expected_retry_indices=indices, ) calibration = _calibration(args, merged) threshold = float(calibration["threshold"]) confidence_key = str(calibration["confidence_key"]) decided = [ apply_gate_decision( record, confidence_key=confidence_key, threshold=threshold, ) for record in merged ] summary = summarize_teacher_gate_records( decided, confidence_key=confidence_key, ) _atomic_write_jsonl(output_path, decided) manifest = { "schema_version": SCHEMA_VERSION, "complete": True, "split": args.split, "input_path": str(Path(args.input).resolve()), "output_path": str(output_path), "record_count": len(decided), "source_row_count": len(raw_rows), "selected_source_indices": ( None if args.merge_selection_cache or len(indices) == len(raw_rows) else indices ), "dataset_fingerprint": chartqa_dataset_fingerprint( raw_rows if args.merge_selection_cache else selected_samples ), "teacher_model": str(Path(args.teacher_model).resolve()), "processor_model": str( Path(args.processor_model or args.teacher_model).resolve() ), "providers": providers, "human_only": bool(args.human_only), "generation": { "do_sample": int(args.num_candidates) > 1, "num_candidates": int(args.num_candidates), "includes_greedy_candidate": True, "sampled_candidates": max(0, int(args.num_candidates) - 1), "temperature": ( float(args.temperature) if int(args.num_candidates) > 1 else None ), "top_p": float(args.top_p) if int(args.num_candidates) > 1 else None, "seed": int(args.seed), "max_new_tokens": int(args.max_new_tokens), "repetition_penalty": float(args.repetition_penalty), "batch_size_per_rank": int(args.batch_size), "world_size": world_size, }, "source_selection": { "mode": str(args.source_selection), "selection_cache": ( str(Path(args.selection_cache).resolve()) if args.selection_cache else None ), "candidate_source_rows": len(candidate_indices), "merged_complete_cache": bool(args.merge_selection_cache), "retry_record_count": retry_record_count, "recovered_hard_pass_rows": recovered_rows, }, "teacher_draft_source": ( { "cache_path": str(Path(args.teacher_draft_cache).resolve()), "providers": list((draft_manifest or {}).get("providers") or []), "dataset_fingerprint": str( (draft_manifest or {}).get("dataset_fingerprint") or "" ), "applied_without_gold_routing": True, } if args.teacher_draft_cache else None ), "decision_source": ( "validation_calibrated_confidence_plus_train_gold_hard_gate" if args.calibration_manifest else "smoke_or_local_calibration" ), "calibration": calibration, "summary": summary, } _atomic_write_json(Path(default_manifest_path(output_path)), manifest) print( "[teacher-gate][complete] " + json.dumps( {"output": str(output_path), "calibration": calibration, "summary": summary}, ensure_ascii=False, sort_keys=True, ), flush=True, ) _barrier(world_size) return 0 if __name__ == "__main__": raise SystemExit(main())