flow_grpo_cxr / analysis_tools /compare_inference_paths.py
zhui711's picture
Upload folder using huggingface_hub
535fb25 verified
Raw
History Blame Contribute Delete
4.37 kB
#!/usr/bin/env python3
"""Compare patched eval-style generation and external OmniGenPipeline on a tiny set.
This script performs full generation and may require GPU memory. It writes only
under the requested output_dir.
"""
from __future__ import annotations
import argparse
import json
import os
import sys
import numpy as np
import torch
from PIL import Image
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--jsonl_path", default="dataset/cxr_radiomics_current_server/test_metadata.jsonl")
parser.add_argument("--output_dir", default="outputs/compare_inference_paths")
parser.add_argument("--model_path", default="Shitao/OmniGen-v1")
parser.add_argument("--lora_path", default=None)
parser.add_argument("--omnigen_code_root", default=os.environ.get("OMNIGEN_CODE_ROOT", "/home/wenting/zr/gen_code"))
parser.add_argument("--num_samples", type=int, default=1)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--device", default="cuda:0")
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)
if args.omnigen_code_root not in sys.path:
sys.path.insert(0, args.omnigen_code_root)
from flow_grpo.omnigen_patch.omnigen_pipeline_with_logprob import pipeline_with_logprob
from scripts.test_omnigen_cxr import get_input_images, get_instruction, load_jsonl
from scripts.train_omnigen import load_omnigen_components, merge_lora_into_base_model, requires_grad, _to_rgb_pil
from OmniGen import OmniGenPipeline
import ml_collections
os.makedirs(args.output_dir, exist_ok=True)
records = load_jsonl(args.jsonl_path)[: args.num_samples]
prompts = [get_instruction(r) for r in records]
input_images = [get_input_images(r) for r in records]
config = ml_collections.ConfigDict()
config.pretrained = ml_collections.ConfigDict()
config.pretrained.model = args.model_path
config.pretrained.vae_path = None
config.activation_checkpointing = False
device = torch.device(args.device)
dtype = torch.bfloat16
model, vae, processor = load_omnigen_components(config, device, dtype)
requires_grad(vae, False)
if args.lora_path:
model = merge_lora_into_base_model(model, args.lora_path, dtype, trainable=False)
model.eval()
torch.manual_seed(args.seed)
patched = pipeline_with_logprob(
model,
vae,
processor,
prompts,
input_images,
height=256,
width=256,
num_inference_steps=50,
guidance_scale=2.5,
img_guidance_scale=2.0,
max_input_image_size=256,
use_img_guidance=True,
use_input_image_size_as_output=False,
dtype=dtype,
output_type="pt",
noise_level=0.0,
sde_type="cps",
)["images"]
pipe = OmniGenPipeline.from_pretrained(args.model_path)
if args.lora_path:
pipe.merge_lora(args.lora_path)
pipe.to(device)
external = pipe(
prompt=prompts if len(prompts) > 1 else prompts[0],
input_images=input_images if len(prompts) > 1 else input_images[0],
height=256,
width=256,
num_inference_steps=50,
guidance_scale=2.5,
img_guidance_scale=2.0,
use_input_image_size_as_output=False,
use_kv_cache=True,
offload_kv_cache=False,
separate_cfg_infer=False,
offload_model=False,
seed=args.seed,
output_type="pt",
)
if isinstance(external, list):
external = torch.stack([torch.from_numpy(np.asarray(img)).permute(2, 0, 1).float() / 255.0 for img in external])
diffs = (patched.detach().cpu().float() - external.detach().cpu().float()).abs()
for i in range(len(records)):
_to_rgb_pil(patched[i]).save(os.path.join(args.output_dir, f"patched_{i}.png"))
_to_rgb_pil(external[i]).save(os.path.join(args.output_dir, f"external_{i}.png"))
report = {"max_abs_diff": float(diffs.max()), "mean_abs_diff": float(diffs.mean()), "shape": list(diffs.shape)}
with open(os.path.join(args.output_dir, "comparison.json"), "w", encoding="utf-8") as handle:
json.dump(report, handle, indent=2)
print(json.dumps(report, indent=2))
if __name__ == "__main__":
main()