flow_grpo_cxr / analysis_tools /trace_inference_batch.py
zhui711's picture
Upload folder using huggingface_hub
535fb25 verified
Raw
History Blame Contribute Delete
3.25 kB
#!/usr/bin/env python3
"""Print a lightweight JSON summary of one standalone inference batch.
Run from the repository root:
python3 analysis_tools/trace_inference_batch.py \
--jsonl_path dataset/cxr_radiomics_current_server/test_metadata.jsonl \
--batch_size 2
This script does not load OmniGen weights or generate images by default. Use
--include_processor to run external OmniGenProcessor preprocessing.
"""
from __future__ import annotations
import argparse
import json
import os
import sys
from typing import Any
def summarize_value(value: Any):
if hasattr(value, "shape"):
return {
"type": type(value).__name__,
"shape": list(value.shape),
"dtype": str(getattr(value, "dtype", None)),
"device": str(getattr(value, "device", None)),
}
if isinstance(value, dict):
return {str(k): summarize_value(v) for k, v in value.items()}
if isinstance(value, (list, tuple)):
return {
"type": type(value).__name__,
"len": len(value),
"items": [summarize_value(v) for v in list(value)[:3]],
}
return {"type": type(value).__name__, "repr": repr(value)[:300]}
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--jsonl_path", default="dataset/cxr_radiomics_current_server/test_metadata.jsonl")
parser.add_argument("--batch_size", type=int, default=2)
parser.add_argument("--model_path", default="Shitao/OmniGen-v1")
parser.add_argument("--omnigen_code_root", default=os.environ.get("OMNIGEN_CODE_ROOT", "/home/wenting/zr/gen_code"))
parser.add_argument("--height", type=int, default=256)
parser.add_argument("--width", type=int, default=256)
parser.add_argument("--include_processor", action="store_true")
args = parser.parse_args()
repo_root = os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))
if repo_root not in sys.path:
sys.path.insert(0, repo_root)
from scripts.test_omnigen_cxr import get_instruction, get_input_images, load_jsonl
records = load_jsonl(args.jsonl_path)
batch = records[: args.batch_size]
prompts = [get_instruction(item) for item in batch]
input_images = [get_input_images(item) for item in batch]
output = {
"jsonl_path": args.jsonl_path,
"num_records": len(records),
"batch_size": len(batch),
"raw_records": summarize_value(batch),
"prompts": summarize_value(prompts),
"input_images": summarize_value(input_images),
}
if args.include_processor:
if args.omnigen_code_root and args.omnigen_code_root not in sys.path:
sys.path.insert(0, args.omnigen_code_root)
from OmniGen import OmniGenProcessor
processor = OmniGenProcessor.from_pretrained(args.model_path)
input_data = processor(
prompts,
input_images,
height=args.height,
width=args.width,
use_img_cfg=True,
separate_cfg_input=False,
use_input_image_size_as_output=False,
)
output["processor"] = summarize_value(input_data)
print(json.dumps(output, indent=2, default=str))
if __name__ == "__main__":
main()