#!/usr/bin/env python3 """Generate a few ChartQA responses and audit the DyME CoT format.""" from __future__ import annotations import argparse import json import sys from pathlib import Path import torch from PIL import Image from transformers import AutoProcessor, LlavaOnevisionForConditionalGeneration _PROJECT_ROOT = Path(__file__).resolve().parent.parent if str(_PROJECT_ROOT) not in sys.path: sys.path.insert(0, str(_PROJECT_ROOT)) from data_utils.chart.data_collector import prepare_chart_rl_data SECTIONS = ("Goal:", "Observation:", "Reasoning:", "Conclusion:", "Answer:") def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--model_path", required=True) parser.add_argument( "--processor_path", default=None, help="Tokenizer/processor source; defaults to --model_path.", ) parser.add_argument("--dataset", required=True) parser.add_argument("--indices", default="0,1,100,1000") parser.add_argument("--max_new_tokens", type=int, default=300) parser.add_argument("--output", default=None) return parser.parse_args() def ordered_sections(text: str) -> bool: positions = [text.find(section) for section in SECTIONS] return all(position >= 0 for position in positions) and positions == sorted(positions) def main() -> int: args = parse_args() indices = [int(value.strip()) for value in args.indices.split(",") if value.strip()] rows = prepare_chart_rl_data(args.dataset) if not indices or min(indices) < 0 or max(indices) >= len(rows): raise ValueError(f"indices must be within [0, {len(rows) - 1}]") model_path = str(Path(args.model_path).resolve()) processor_path = str(Path(args.processor_path or args.model_path).resolve()) processor = AutoProcessor.from_pretrained(processor_path, local_files_only=True) processor.tokenizer.padding_side = "left" model = LlavaOnevisionForConditionalGeneration.from_pretrained( model_path, torch_dtype=torch.bfloat16, low_cpu_mem_usage=True, attn_implementation="sdpa", local_files_only=True, ).to("cuda:0") model.eval() results = [] for index in indices: row = rows[index] image = Image.open(row["image"]).convert("RGB") messages = [ { "role": "user", "content": [ {"type": "image"}, {"type": "text", "text": row["prompt"]}, ], } ] text = processor.apply_chat_template( messages, add_generation_prompt=True, tokenize=False ) inputs = processor(text=[text], images=[image], return_tensors="pt") inputs = {key: value.to("cuda:0") for key, value in inputs.items()} prompt_length = inputs["input_ids"].shape[1] with torch.inference_mode(): output_ids = model.generate( **inputs, max_new_tokens=args.max_new_tokens, do_sample=False, repetition_penalty=1.0, use_cache=True, ) generated_ids = output_ids[:, prompt_length:] response = processor.batch_decode( generated_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False )[0].strip() result = { "index": index, "question": row["question_wo_prompt"], "gold": row["reference_answer"], "response": response, "ordered_cot_format": ordered_sections(response), "response_words": len(response.split()), } results.append(result) print( f"\n[COT sample {index}] format={result['ordered_cot_format']} " f"words={result['response_words']} gold={result['gold']}\n{response}", flush=True, ) format_count = sum(result["ordered_cot_format"] for result in results) summary = { "model_path": model_path, "processor_path": processor_path, "samples": len(results), "ordered_cot_count": format_count, "ordered_cot_rate": format_count / len(results), "mean_response_words": sum(item["response_words"] for item in results) / len(results), "results": results, } print("\n[COT summary] " + json.dumps({ key: value for key, value in summary.items() if key != "results" }, ensure_ascii=False), flush=True) if args.output: output_path = Path(args.output) output_path.parent.mkdir(parents=True, exist_ok=True) output_path.write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding="utf-8") return 0 if __name__ == "__main__": raise SystemExit(main())