omnigen_seg / debug_dataloader.py
zhui711's picture
Upload folder using huggingface_hub
4f08626 verified
Raw
History Blame Contribute Delete
3.52 kB
#!/usr/bin/env python3
"""
Debug script to inspect CXRSegDataset output.
Checks whether the Heart channel (idx 2, class_id 51) is all-zeros.
"""
import sys
import os
import argparse
# Make the vendored smp package importable
sys.path.insert(0, os.path.join(os.path.dirname(__file__),
"segmentation_models_pytorch"))
from segmentation_models_pytorch.datasets.cxr_seg_dataset import CXRSegDataset
def main():
parser = argparse.ArgumentParser(
description="Inspect CXRSegDataset mask channels for empty Heart/Aorta masks")
parser.add_argument("--data_dir", type=str, required=True,
help="Path to lidc_TotalSeg directory")
parser.add_argument("--num_samples", type=int, default=50,
help="Number of samples to inspect (default: 50)")
args = parser.parse_args()
# Channel name lookup (matches TARGET_GROUPS insertion order)
channel_names = list(CXRSegDataset.TARGET_GROUPS.keys())
HEART_CH = 2 # "Heart"
AORTA_CH = 3 # "Aorta"
print(f"Channel mapping: {dict(enumerate(channel_names))}")
print(f"Heart -> channel {HEART_CH} (class IDs {CXRSegDataset.TARGET_GROUPS['Heart']})")
print(f"Aorta -> channel {AORTA_CH} (class IDs {CXRSegDataset.TARGET_GROUPS['Aorta']})")
print()
# Instantiate dataset without albumentations — uses fallback tensor conversion
dataset = CXRSegDataset(root_dir=args.data_dir, transform=None)
total = len(dataset)
print(f"Dataset contains {total} samples total.")
n = min(args.num_samples, total)
if n == 0:
print("No samples found. Exiting.")
return
print(f"Inspecting the first {n} samples ...\n")
empty_heart = 0
empty_aorta = 0
for idx in range(n):
batch = dataset[idx]
mask = batch["mask"] # shape: (10, 256, 256)
sums = [mask[ch].sum().item() for ch in range(mask.shape[0])]
heart_sum = sums[HEART_CH]
aorta_sum = sums[AORTA_CH]
if heart_sum == 0:
empty_heart += 1
if aorta_sum == 0:
empty_aorta += 1
# Per-sample details for the first 10 and any with empty Heart
if idx < 10 or heart_sum == 0:
tag = " *** EMPTY Heart ***" if heart_sum == 0 else ""
sample_info = batch.get("_debug_path", dataset.samples[idx]["mask_path"])
print(f"[{idx:3d}] {os.path.basename(sample_info)}{tag}")
for ch, name in enumerate(channel_names):
print(f" ch{ch} {name:16s}: sum={sums[ch]:10.1f}")
print()
# ---- Summary ----
print("=" * 60)
print(f"SUMMARY (first {n} samples)")
print(f" Empty Heart masks : {empty_heart} / {n}")
print(f" Empty Aorta masks : {empty_aorta} / {n}")
print()
if empty_heart == n:
print("!! ALL Heart masks are empty — likely a data-loading bug.")
print(" Check: is class_id 51 present in labels_found for these patients?")
# Quick diagnosis: check labels_found for first sample
lf = dataset.samples[0]["labels_found"]
print(f" Sample 0 labels_found includes 51? -> {51 in lf}")
print(f" Sample 0 labels_found (first 20): {sorted(lf)[:20]}")
elif empty_heart > 0:
print(f"!! {empty_heart} samples have empty Heart masks (partial issue).")
else:
print("All samples have non-empty Heart masks. Channel 2 looks fine.")
print("=" * 60)
if __name__ == "__main__":
main()