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