umm / scripts /trace_cvbench_cases.py
Orangerl's picture
Update Stage-2 code, evaluations, handoff, and deployment skill
bca45b0 verified
Raw History Blame Contribute Delete
9.86 kB
#!/usr/bin/env python3
"""Capture observable token-by-token generation traces for selected CV-Bench cases."""
from __future__ import annotations
import argparse
import ctypes
import json
from datetime import datetime, timezone
from pathlib import Path
import torch
from transformers import AutoProcessor
from cvbench_eval import CVBenchEvalCallback, _expected_choice, extract_choice
from models.blip3o.model.language_model.covt_qwen_stage2_van import (
CoVTVanForConditionalGeneration,
)
ANCHORS = ["sam", "dino", "depth", "pidinet", "siglip"]
COVT_INSTRUCTION = (
"When given a question: {Question} and its corresponding image,"
"you need to output your reasoning process within <think></think> tags and provide "
"the final answer within <answer></answer> tags."
"The reasoning process should include some visual chain-of-thought tokens, such as "
"<|dino_pad|>, <|depth_pad|>, <|pidinet_pad|>, and <|siglip_pad|>. You must adhere "
"to this format when producing the output."
"i.e., <think> thinking process here </think>"
"<answer>...</answer>"
)
def _set_process_name(name: str) -> None:
try:
ctypes.CDLL(None).prctl(15, name.encode()[:15], 0, 0, 0)
except Exception:
pass
def _candidate(tokenizer, token_id: int, probability: float) -> dict:
return {
"token_id": token_id,
"token": tokenizer.convert_ids_to_tokens(token_id),
"text_piece": tokenizer.decode(
[token_id], skip_special_tokens=False, clean_up_tokenization_spaces=False
),
"probability": probability,
}
def _generate_trace(model, processor, inputs: dict, *, max_new_tokens: int, top_k: int) -> dict:
device = next(model.parameters()).device
dtype = next(model.parameters()).dtype
inputs = {key: value.to(device) for key, value in inputs.items()}
if "pixel_values" in inputs:
inputs["pixel_values"] = inputs["pixel_values"].to(dtype=dtype)
if "pixel_values_videos" in inputs:
inputs["pixel_values_videos"] = inputs["pixel_values_videos"].to(dtype=dtype)
generation_inputs = {
"input_ids": inputs["input_ids"],
"attention_mask": inputs.get("attention_mask"),
"max_new_tokens": max_new_tokens,
"do_sample": False,
"use_cache": True,
"temperature": 1.0,
"top_p": 1.0,
"top_k": 50,
"pad_token_id": processor.tokenizer.pad_token_id
or processor.tokenizer.eos_token_id,
"eos_token_id": processor.tokenizer.eos_token_id,
"return_dict_in_generate": True,
"output_scores": True,
}
for key in (
"pixel_values",
"image_grid_thw",
"pixel_values_videos",
"video_grid_thw",
"mm_token_type_ids",
):
if key in inputs:
generation_inputs[key] = inputs[key]
previous_rope_deltas = model.rope_deltas
model.rope_deltas = None
try:
output = model.generate(**generation_inputs)
finally:
model.rope_deltas = previous_rope_deltas
prompt_length = int(inputs["input_ids"].shape[1])
generated_ids = output.sequences[0, prompt_length:]
steps = []
for position, (token_id_tensor, score_tensor) in enumerate(
zip(generated_ids, output.scores), start=1
):
token_id = int(token_id_tensor.item())
logits = score_tensor[0].float()
log_denominator = torch.logsumexp(logits, dim=-1)
selected_probability = float(torch.exp(logits[token_id] - log_denominator).item())
top_values, top_ids = torch.topk(logits, k=top_k)
candidates = [
_candidate(
processor.tokenizer,
int(candidate_id.item()),
float(torch.exp(value - log_denominator).item()),
)
for value, candidate_id in zip(top_values, top_ids)
]
selected = _candidate(processor.tokenizer, token_id, selected_probability)
steps.append({"position": position, "selected": selected, "top_candidates": candidates})
raw_text = processor.tokenizer.decode(
generated_ids,
skip_special_tokens=False,
clean_up_tokenization_spaces=False,
)
visible_text = processor.tokenizer.decode(
generated_ids,
skip_special_tokens=True,
clean_up_tokenization_spaces=False,
)
return {
"prompt_token_count": prompt_length,
"generated_token_count": len(steps),
"generated_token_ids": [int(value) for value in generated_ids.tolist()],
"raw_text_with_special_tokens": raw_text,
"visible_text_skip_special_tokens": visible_text,
"steps": steps,
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--model-path", type=Path, required=True)
parser.add_argument("--manifest", type=Path, required=True)
parser.add_argument("--sample-ids", nargs="+", required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--max-new-tokens", type=int, default=64)
parser.add_argument("--top-k", type=int, default=5)
parser.add_argument(
"--modes",
nargs="+",
choices=("faithful_eval", "training_style", "covt_required"),
default=("faithful_eval", "training_style"),
)
args = parser.parse_args()
_set_process_name("H3_TRACE")
model_path = args.model_path.resolve()
manifest_path = args.manifest.resolve()
output_path = args.output.resolve()
records = [
json.loads(line)
for line in manifest_path.read_text(encoding="utf-8").splitlines()
if line
]
by_id = {record["sample_id"]: record for record in records}
missing = [sample_id for sample_id in args.sample_ids if sample_id not in by_id]
if missing:
raise KeyError(f"Unknown sample IDs: {missing}")
processor = AutoProcessor.from_pretrained(
model_path, local_files_only=True, trust_remote_code=True
)
model = CoVTVanForConditionalGeneration.from_pretrained(
model_path,
local_files_only=True,
torch_dtype=torch.bfloat16,
attn_implementation="flash_attention_2",
low_cpu_mem_usage=True,
device_map={"": 0},
)
model.get_anchor_model_ids(ANCHORS, load_anchor_models=False)
model.config.use_cache = True
model.eval()
evaluator = CVBenchEvalCallback(
processor=processor,
manifest_path=str(manifest_path),
every_n_steps=1,
max_new_tokens=args.max_new_tokens,
fail_fast=True,
)
evaluator._load_records()
results = []
with torch.inference_mode():
for sample_id in args.sample_ids:
record = next(row for row in evaluator._records if row["sample_id"] == sample_id)
for mode in args.modes:
if mode == "faithful_eval":
prompt = (
str(record["prompt"])
+ "\nReturn the selected option in <answer> tags, for example "
+ "<answer>(A)</answer>."
)
elif mode == "covt_required":
# Keep this byte-for-byte aligned with CoVTVan.tokenize_fn's
# intended inference wrapper so special CoVT tokens can be audited.
prompt = COVT_INSTRUCTION.format(Question=str(record["prompt"]))
else:
prompt = str(record["prompt"])
from PIL import Image
with Image.open(record["_image_path"]) as opened:
inputs = evaluator._tokenize(prompt, opened.convert("RGB"))
trace = _generate_trace(
model,
processor,
inputs,
max_new_tokens=args.max_new_tokens,
top_k=args.top_k,
)
predicted = extract_choice(
trace["visible_text_skip_special_tokens"], record["choices"]
)
expected = _expected_choice(record["answer"])
results.append(
{
"sample_id": sample_id,
"config": record["config"],
"task": record["task"],
"image": record["image"],
"question": record["question"],
"choices": record["choices"],
"expected": expected,
"mode": mode,
"prompt": prompt,
"predicted": predicted,
"correct": predicted == expected,
"trace": trace,
}
)
print(
f"[trace] {sample_id}/{mode}: "
f"raw={trace['raw_text_with_special_tokens']!r} "
f"predicted={predicted} expected={expected}",
flush=True,
)
payload = {
"created_utc": datetime.now(timezone.utc).isoformat(),
"model_path": str(model_path),
"manifest": str(manifest_path),
"dtype": "bfloat16",
"attention": "flash_attention_2",
"decoding": "greedy",
"max_new_tokens": args.max_new_tokens,
"top_k": args.top_k,
"scope_note": (
"Observable emitted-token trace only. It does not expose or claim to reconstruct "
"the model's hidden-state reasoning process."
),
"results": results,
}
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_text(
json.dumps(payload, ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
)
print(f"[trace] saved {output_path}")
if __name__ == "__main__":
main()