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