#!/usr/bin/env python3 """Print a lightweight JSON summary of one CXR training batch. Run from the repository root: python3 analysis_tools/trace_training_batch.py \ --config config/grpo.py:general_radiomics_omnigen_4gpu_kl \ --split train --batch_size 2 This script intentionally does not load OmniGen weights or generate images by default. Use --include_processor to also instantiate OmniGenProcessor and run the text/image preprocessing path for the sampled batch. """ from __future__ import annotations import argparse import importlib.util import json import os import sys from typing import Any from torch.utils.data import DataLoader def load_config(entry: str): path, name = entry.split(":", 1) spec = importlib.util.spec_from_file_location("trace_config", path) if spec is None or spec.loader is None: raise ImportError(f"Could not load config module: {path}") module = importlib.util.module_from_spec(spec) spec.loader.exec_module(module) return getattr(module, name)() 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("--config", default="config/grpo.py:general_radiomics_omnigen_4gpu_kl") parser.add_argument("--split", default="train", choices=["train", "test"]) parser.add_argument("--batch_size", type=int, default=2) 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) config = load_config(args.config) from scripts.train_omnigen import RadiomicsEditDataset dataset = RadiomicsEditDataset(config.dataset, args.split, condition_dropout_prob=0.0) loader = DataLoader( dataset, batch_size=args.batch_size, shuffle=False, collate_fn=RadiomicsEditDataset.collate_fn, ) prompts, instructions, metadatas, input_image_paths, ref_images, group_keys = next(iter(loader)) output = { "config_entry": args.config, "dataset_root": config.dataset, "split": args.split, "dataset_len": len(dataset), "batch": { "prompts": summarize_value(prompts), "instructions": summarize_value(instructions), "metadatas": summarize_value(metadatas), "input_image_paths": summarize_value(input_image_paths), "ref_images": [img.size for img in ref_images], "group_keys": summarize_value(group_keys), }, } if args.include_processor: omnigen_root = os.environ.get("OMNIGEN_CODE_ROOT", getattr(config.pretrained, "local_code_root", "")) if omnigen_root and omnigen_root not in sys.path: sys.path.insert(0, omnigen_root) from OmniGen import OmniGenProcessor from scripts.train_omnigen import resolve_model_root model_root = resolve_model_root(config.pretrained.model) processor = OmniGenProcessor.from_pretrained(model_root) input_data = processor( list(instructions), list(input_image_paths), height=config.resolution, width=config.resolution, use_img_cfg=config.sample.use_img_guidance, separate_cfg_input=True, use_input_image_size_as_output=config.sample.use_input_image_size_as_output, ) output["processor"] = summarize_value(input_data) print(json.dumps(output, indent=2, default=str)) if __name__ == "__main__": main()