Download scripts/preprocess_chartqa_teacher_gate.py from Jack04810/agentic-rl-main: direct link, hf CLI and curl.
- Browser
- Download file 43.3 kB
-
https://huggingface.co/Jack04810/agentic-rl-main/resolve/main/scripts/preprocess_chartqa_teacher_gate.py
- Command line
-
hf download hf://Jack04810/agentic-rl-main/scripts/preprocess_chartqa_teacher_gate.py
-
curl -L -o preprocess_chartqa_teacher_gate.py https://huggingface.co/Jack04810/agentic-rl-main/resolve/main/scripts/preprocess_chartqa_teacher_gate.py
43.3 kB
| #!/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()) | |