#!/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()