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