| |
| """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() |
|
|
|
|