agentic-rl-main / scripts /check_chartqa_cot_generation.py
Jack04810's picture
Add files using upload-large-folder tool
df529cc verified
Raw History Blame Contribute Delete
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())