File size: 4,235 Bytes
535fb25 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 | #!/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()
|