""" pre_process.py -- Part A, Section 2.1 (Data preparation and exploratory analysis) What this script does --------------------- 1. Indexes the COVID-19 Radiography Database *without modifying any original file*. 2. Runs a data-quality audit: missing masks, unreadable files, duplicate images (MD5 of file bytes + perceptual hash of pixels), unusual image sizes, image/mask size mismatches, and mask-area statistics. 3. Creates frozen, stratified train/val/test splits with SEED = 16 and writes splits/covid_classification_seed16.csv and splits/covid_segmentation_seed16.csv. 4. Produces the exploratory plots (class distribution, resolution distribution, intensity summaries, mask-area distribution) and the seed-selected qualitative figure (original / preprocessed / mask overlay, with image identifiers). Usage ----- python pre_process.py --data-root ../Dataset/covid19/COVID-19_Radiography_Dataset python pre_process.py # auto-detects under ../Dataset/covid19 python pre_process.py --skip-hash # faster audit, no duplicate detection """ from __future__ import annotations import argparse import hashlib from collections import Counter, defaultdict from pathlib import Path import numpy as np import pandas as pd from common import (CLASS_DIR_ALIASES, boxplot, CLASSES, CLS_SPLIT_CSV, DATASET_DIR, IMG_SIZE, METRICS_DIR, OVERLAY_DIR, PLOTS_DIR, SEED, SEED_TAG, SEG_SPLIT_CSV, SPLIT_FRACTIONS, _mpl, file_md5, load_image_gray, load_mask_bin, overlay_mask_on_image, progress, save_json, seeded_rng, set_seed) IMG_EXT = {".png", ".jpg", ".jpeg", ".bmp", ".tif", ".tiff"} KAGGLE_SLUG = "tawsifurrahman/covid19-radiography-database" # --------------------------------------------------------------------------- # Dataset discovery (local folder, kagglehub cache, or kagglehub download) # --------------------------------------------------------------------------- def _kagglehub_path(slug: str, allow_download: bool = True) -> Path | None: """Return the local kagglehub cache path for `slug`, downloading if needed.""" try: import kagglehub except ImportError: print("[info] kagglehub is not installed (pip install kagglehub) -- " "falling back to a local Dataset/ folder") return None try: if not allow_download: return None print(f"[info] resolving '{slug}' through kagglehub (this downloads once, " f"then reuses the cache) ...") return Path(kagglehub.dataset_download(slug)) except Exception as e: print(f"[warn] kagglehub could not provide '{slug}': {e}") return None def find_data_root(user_root: str | None, use_kagglehub: bool = True) -> Path: """Locate the folder that directly contains the four class folders. Search order: 1. --data-root, if given 2. Dataset/covid19/... inside this submission folder 3. the kagglehub cache (downloading the dataset on first use) """ candidates: list[Path] = [] if user_root: candidates.append(Path(user_root)) candidates += [DATASET_DIR / "COVID-19_Radiography_Dataset", DATASET_DIR] if DATASET_DIR.exists(): candidates += [p for p in DATASET_DIR.rglob("COVID-19_Radiography_Dataset") if p.is_dir()] for c in candidates: if c and c.exists() and _resolve_class_dirs(c): return c.resolve() if use_kagglehub: kp = _kagglehub_path(KAGGLE_SLUG) if kp is not None and kp.exists(): for c in [kp / "COVID-19_Radiography_Dataset", kp, *[p for p in kp.rglob("COVID-19_Radiography_Dataset") if p.is_dir()]]: if _resolve_class_dirs(c): _link_into_dataset_dir(c) return c.resolve() raise FileNotFoundError( "Could not find the COVID-19 Radiography Database.\n" f"Looked under: {DATASET_DIR} and the kagglehub cache.\n" "Either run python download_data.py from the submission root, or unzip the " "Kaggle download so that the class folders (COVID/, Lung_Opacity/, Normal/, " f"'Viral Pneumonia'/) live under {DATASET_DIR}/COVID-19_Radiography_Dataset/, " "or pass --data-root explicitly." ) def _link_into_dataset_dir(src: Path) -> None: """Point Dataset/covid19/COVID-19_Radiography_Dataset at the kagglehub copy so the submitted folder structure is self-describing without duplicating ~800 MB.""" dst = DATASET_DIR / "COVID-19_Radiography_Dataset" try: DATASET_DIR.mkdir(parents=True, exist_ok=True) if dst.exists() or dst.is_symlink(): return dst.symlink_to(src, target_is_directory=True) print(f"[info] linked {dst} -> {src}") except Exception: pass # symlinks may be unavailable (e.g. Windows without dev mode) def _resolve_class_dirs(root: Path) -> dict[str, Path]: """Map canonical class name -> its folder on disk (empty dict if not a valid root).""" found = {} for canon, aliases in CLASS_DIR_ALIASES.items(): for a in aliases: p = root / a if p.is_dir(): found[canon] = p break return found if len(found) == len(CLASSES) else {} def _images_and_masks(class_dir: Path) -> tuple[list[Path], Path | None]: """Handle both dataset layouts: /images + /masks (v3+), or a flat / folder of images (older mirrors).""" img_dir = class_dir / "images" if (class_dir / "images").is_dir() else class_dir mask_dir = class_dir / "masks" if (class_dir / "masks").is_dir() else None imgs = sorted(p for p in img_dir.iterdir() if p.suffix.lower() in IMG_EXT) return imgs, mask_dir def _match_mask(mask_dir: Path | None, stem: str) -> Path | None: if mask_dir is None: return None for ext in (".png", ".jpg", ".jpeg", ".tif", ".bmp"): p = mask_dir / f"{stem}{ext}" if p.exists(): return p return None # --------------------------------------------------------------------------- # Indexing + quality audit # --------------------------------------------------------------------------- def phash(path: Path, size: int = 8) -> str: """Tiny perceptual hash (average hash) -- catches near-duplicates that have been re-encoded, which a byte-level MD5 would miss.""" from PIL import Image im = Image.open(path).convert("L").resize((size, size), Image.BILINEAR) a = np.asarray(im, dtype=np.float32) bits = (a > a.mean()).flatten() return "".join("1" if b else "0" for b in bits) def build_index(root: Path, compute_hashes: bool = True) -> tuple[pd.DataFrame, dict]: class_dirs = _resolve_class_dirs(root) rows, issues = [], defaultdict(list) for cls in CLASSES: cdir = class_dirs[cls] imgs, mask_dir = _images_and_masks(cdir) if mask_dir is None: issues["classes_without_mask_folder"].append(cls) for ip in progress(imgs, desc=f"index {cls}", total=len(imgs)): rec = { "image_id": ip.stem, "label": cls, "image_path": str(ip.resolve()), "mask_path": "", "has_mask": False, "img_w": np.nan, "img_h": np.nan, "img_mode": "", "mask_w": np.nan, "mask_h": np.nan, "mean_intensity": np.nan, "std_intensity": np.nan, "p01": np.nan, "p99": np.nan, "mask_area_frac": np.nan, "md5": "", "phash": "", "readable": True, "mask_readable": True, } # --- image: read the file ONCE and derive everything from it ------ try: from io import BytesIO from PIL import Image raw = ip.read_bytes() if compute_hashes: rec["md5"] = hashlib.md5(raw).hexdigest() with Image.open(BytesIO(raw)) as im: rec["img_w"], rec["img_h"] = im.size rec["img_mode"] = im.mode gray = im.convert("L") a = np.asarray(gray, dtype=np.float32) / 255.0 if compute_hashes: small = np.asarray(gray.resize((8, 8), Image.BILINEAR), dtype=np.float32) rec["phash"] = "".join( "1" if b else "0" for b in (small > small.mean()).flatten()) rec["mean_intensity"] = float(a.mean()) rec["std_intensity"] = float(a.std()) rec["p01"], rec["p99"] = float(np.percentile(a, 1)), float(np.percentile(a, 99)) except Exception as e: rec["readable"] = False issues["unreadable_images"].append({"path": str(ip), "error": str(e)}) # --- mask --- mp = _match_mask(mask_dir, ip.stem) if mp is None: issues["missing_masks"].append({"image_id": ip.stem, "label": cls}) else: rec["mask_path"] = str(mp.resolve()) try: from PIL import Image with Image.open(mp) as mm: rec["mask_w"], rec["mask_h"] = mm.size m = np.asarray(mm.convert("L"), dtype=np.float32) / 255.0 b = (m > 0.5) rec["mask_area_frac"] = float(b.mean()) rec["has_mask"] = True if b.sum() == 0: issues["empty_masks"].append({"image_id": ip.stem, "label": cls}) uniq = np.unique((m * 255).astype(np.uint8)) if len(uniq) > 2 and not set(uniq.tolist()) <= {0, 255}: issues["non_binary_masks"].append( {"image_id": ip.stem, "n_levels": int(len(uniq))}) except Exception as e: rec["mask_readable"] = False issues["unreadable_masks"].append({"path": str(mp), "error": str(e)}) if rec["readable"] and (rec["img_w"], rec["img_h"]) != (rec["mask_w"], rec["mask_h"]) \ and rec["has_mask"]: issues["image_mask_size_mismatch"].append( {"image_id": ip.stem, "image_size": [rec["img_w"], rec["img_h"]], "mask_size": [rec["mask_w"], rec["mask_h"]]}) rows.append(rec) df = pd.DataFrame(rows) # ---- duplicates ---- if compute_hashes and df["md5"].str.len().gt(0).any(): dup_md5 = df[df["md5"].duplicated(keep=False) & df["md5"].ne("")] for h, g in dup_md5.groupby("md5"): issues["duplicate_images_exact"].append( {"md5": h, "image_ids": g["image_id"].tolist(), "labels": g["label"].tolist()}) dup_ph = df[df["phash"].duplicated(keep=False) & df["phash"].ne("")] for h, g in dup_ph.groupby("phash"): if len(g) > 1 and g["md5"].nunique() > 1: issues["duplicate_images_perceptual"].append( {"phash": h, "image_ids": g["image_id"].tolist(), "labels": g["label"].tolist()}) # ---- unusual sizes ---- sizes = Counter(zip(df["img_w"].tolist(), df["img_h"].tolist())) if sizes: modal = sizes.most_common(1)[0][0] odd = df[(df["img_w"] != modal[0]) | (df["img_h"] != modal[1])] issues["modal_image_size"] = [int(modal[0]), int(modal[1])] issues["unusual_image_sizes"] = ( odd[["image_id", "label", "img_w", "img_h"]].to_dict("records")[:200]) issues["n_unusual_image_sizes"] = int(len(odd)) non_l = df[~df["img_mode"].isin(["L", "1"])] issues["non_grayscale_modes"] = Counter(non_l["img_mode"].tolist()) issues["n_non_grayscale"] = int(len(non_l)) return df, dict(issues) # --------------------------------------------------------------------------- # Splits # --------------------------------------------------------------------------- def make_splits(df: pd.DataFrame, seed: int = SEED, fracs=SPLIT_FRACTIONS) -> pd.DataFrame: """Stratified train/val/test. Duplicate groups (identical MD5) are kept in the same split so an exact copy of a training image can never appear in test.""" from sklearn.model_selection import train_test_split d = df[df["readable"]].copy().reset_index(drop=True) # group key: md5 when available, otherwise the unique image id d["group"] = np.where(d["md5"].astype(str).str.len() > 0, d["md5"], d["image_id"]) grp = (d.groupby("group") .agg(label=("label", "first"), n=("image_id", "size")) .reset_index()) tr, tmp = train_test_split( grp, train_size=fracs[0], random_state=seed, stratify=grp["label"], shuffle=True) rel_val = fracs[1] / (fracs[1] + fracs[2]) va, te = train_test_split( tmp, train_size=rel_val, random_state=seed, stratify=tmp["label"], shuffle=True) assign = {} for name, part in (("train", tr), ("val", va), ("test", te)): for g in part["group"]: assign[g] = name d["split"] = d["group"].map(assign) d = d.drop(columns=["group"]) return d def report_split_counts(d: pd.DataFrame) -> dict: tab = (d.pivot_table(index="label", columns="split", values="image_id", aggfunc="count", fill_value=0) .reindex(index=CLASSES, columns=["train", "val", "test"], fill_value=0)) tab["total"] = tab.sum(axis=1) mask_tab = (d[d["has_mask"]] .pivot_table(index="label", columns="split", values="image_id", aggfunc="count", fill_value=0) .reindex(index=CLASSES, columns=["train", "val", "test"], fill_value=0)) print("\n=== Classification split counts (seed=%d) ===" % SEED) print(tab.to_string()) print("\n=== Images WITH a lung mask (segmentation pool) ===") print(mask_tab.to_string()) return {"classification": tab.to_dict(), "segmentation": mask_tab.to_dict()} # --------------------------------------------------------------------------- # Exploratory plots # --------------------------------------------------------------------------- def eda_plots(d: pd.DataFrame, out_dir: Path = PLOTS_DIR) -> dict: plt = _mpl() out_dir.mkdir(parents=True, exist_ok=True) stats = {} # 1. class distribution per split fig, ax = plt.subplots(figsize=(8, 4.5)) piv = (d.pivot_table(index="label", columns="split", values="image_id", aggfunc="count", fill_value=0) .reindex(index=CLASSES, columns=["train", "val", "test"], fill_value=0)) piv.plot(kind="bar", stacked=True, ax=ax, colormap="viridis") ax.set_title(f"COVID-19 Radiography: class distribution per split (seed={SEED})") ax.set_ylabel("images") ax.set_xlabel("") for i, tot in enumerate(piv.sum(axis=1)): ax.text(i, tot, f" {int(tot)}", ha="center", va="bottom", fontsize=9) plt.setp(ax.get_xticklabels(), rotation=20, ha="right") fig.tight_layout() fig.savefig(out_dir / f"A1_class_distribution_{SEED_TAG}.png", dpi=160) plt.close(fig) stats["class_counts"] = piv.to_dict() counts = piv.sum(axis=1) stats["imbalance_ratio_max_over_min"] = float(counts.max() / max(1, counts.min())) # 2. resolution distribution fig, ax = plt.subplots(figsize=(7, 4.2)) res = d.groupby(["img_w", "img_h"]).size().sort_values(ascending=False).head(12) labels = [f"{int(w)}x{int(h)}" for (w, h) in res.index] ax.bar(labels, res.values, color="#4C72B0") ax.set_yscale("log") ax.set_title("Image resolution distribution (log scale)") ax.set_ylabel("images") plt.setp(ax.get_xticklabels(), rotation=35, ha="right") fig.tight_layout() fig.savefig(out_dir / f"A2_resolution_distribution_{SEED_TAG}.png", dpi=160) plt.close(fig) stats["resolutions"] = {f"{int(w)}x{int(h)}": int(v) for (w, h), v in res.items()} # 3. intensity summaries per class fig, axes = plt.subplots(1, 2, figsize=(12, 4.4)) for c in CLASSES: sub = d[d["label"] == c]["mean_intensity"].dropna() axes[0].hist(sub, bins=60, alpha=0.5, label=c, density=True) axes[0].set_title("Mean image intensity by class") axes[0].set_xlabel("mean intensity (0-1)") axes[0].legend(fontsize=8) data = [d[d["label"] == c]["std_intensity"].dropna().values for c in CLASSES] boxplot(axes[1], data, CLASSES, showfliers=False) axes[1].set_title("Intensity standard deviation by class") plt.setp(axes[1].get_xticklabels(), rotation=20, ha="right") fig.tight_layout() fig.savefig(out_dir / f"A3_intensity_summary_{SEED_TAG}.png", dpi=160) plt.close(fig) stats["intensity_by_class"] = { c: {"mean": float(d[d.label == c].mean_intensity.mean()), "std": float(d[d.label == c].mean_intensity.std()), "p01": float(d[d.label == c].p01.mean()), "p99": float(d[d.label == c].p99.mean())} for c in CLASSES} # 4. mask area distribution md = d[d["has_mask"]] if len(md): fig, axes = plt.subplots(1, 2, figsize=(12, 4.4)) axes[0].hist(md["mask_area_frac"].dropna(), bins=60, color="#C44E52") axes[0].set_title("Lung-mask area as a fraction of the image") axes[0].set_xlabel("mask area fraction") boxplot(axes[1], [md[md.label == c]["mask_area_frac"].dropna().values for c in CLASSES], CLASSES, showfliers=False) axes[1].set_title("Mask area fraction by class") plt.setp(axes[1].get_xticklabels(), rotation=20, ha="right") fig.tight_layout() fig.savefig(out_dir / f"A4_mask_area_distribution_{SEED_TAG}.png", dpi=160) plt.close(fig) stats["mask_area_fraction"] = { "overall": {"mean": float(md.mask_area_frac.mean()), "std": float(md.mask_area_frac.std()), "min": float(md.mask_area_frac.min()), "p05": float(md.mask_area_frac.quantile(0.05)), "median": float(md.mask_area_frac.median()), "p95": float(md.mask_area_frac.quantile(0.95)), "max": float(md.mask_area_frac.max())}, "by_class": {c: float(md[md.label == c].mask_area_frac.mean()) for c in CLASSES}} return stats def qualitative_examples(d: pd.DataFrame, n: int = 8, out_dir: Path = OVERLAY_DIR): """Seed-selected examples: original | preprocessed (224, [0,1]) | mask overlay. Image identifiers are printed on every panel so the report is auditable.""" plt = _mpl() rng = seeded_rng(offset=1) pool = d[d["has_mask"]] picks = [] per_class = max(1, n // len(CLASSES)) for c in CLASSES: sub = pool[pool["label"] == c] if len(sub) == 0: continue idx = rng.choice(len(sub), size=min(per_class, len(sub)), replace=False) picks.append(sub.iloc[idx]) sel = pd.concat(picks).head(n) if picks else pool.head(n) fig, axes = plt.subplots(len(sel), 3, figsize=(9.5, 3.1 * len(sel))) if len(sel) == 1: axes = axes[None, :] from PIL import Image for r, (_, row) in enumerate(sel.iterrows()): orig = np.asarray(Image.open(row["image_path"]).convert("L"), np.float32) / 255.0 prep = load_image_gray(row["image_path"], IMG_SIZE) mask = load_mask_bin(row["mask_path"], IMG_SIZE) axes[r, 0].imshow(orig, cmap="gray") axes[r, 0].set_title(f"{row['image_id']}\noriginal {orig.shape[1]}x{orig.shape[0]} | {row['label']}", fontsize=8) axes[r, 1].imshow(prep, cmap="gray") axes[r, 1].set_title(f"preprocessed {IMG_SIZE}x{IMG_SIZE}, [0,1]", fontsize=8) axes[r, 2].imshow(overlay_mask_on_image(prep, mask)) axes[r, 2].set_title(f"lung mask overlay (area={mask.mean():.3f})", fontsize=8) for c in range(3): axes[r, c].axis("off") fig.suptitle(f"Seed-selected examples (seed={SEED})", fontsize=12) fig.tight_layout() out = out_dir / f"A5_seed_selected_examples_{SEED_TAG}.png" fig.savefig(out, dpi=150) plt.close(fig) print(f"[saved] {out}") return sel[["image_id", "label", "split", "img_w", "img_h", "mask_area_frac"]] # --------------------------------------------------------------------------- # Main # --------------------------------------------------------------------------- def main(): ap = argparse.ArgumentParser(description="Part A data prep + EDA (seed 16)") ap.add_argument("--data-root", default=None, help="folder containing COVID/, Lung_Opacity/, Normal/, 'Viral Pneumonia'/") ap.add_argument("--skip-hash", action="store_true", help="skip MD5/pHash (faster, but no duplicate detection)") ap.add_argument("--n-examples", type=int, default=8) args = ap.parse_args() set_seed(SEED) root = find_data_root(args.data_root) print(f"[info] dataset root : {root}") print(f"[info] seed : {SEED}") df, issues = build_index(root, compute_hashes=not args.skip_hash) print(f"[info] indexed {len(df)} images") d = make_splits(df, seed=SEED) counts = report_split_counts(d) # ---- write the two frozen split files ------------------------------- cls_cols = ["image_id", "label", "split", "image_path", "mask_path", "has_mask", "img_w", "img_h", "mean_intensity", "std_intensity", "mask_area_frac", "md5"] d[cls_cols].to_csv(CLS_SPLIT_CSV, index=False) print(f"[saved] {CLS_SPLIT_CSV}") seg = d[d["has_mask"]].copy() seg[cls_cols].to_csv(SEG_SPLIT_CSV, index=False) print(f"[saved] {SEG_SPLIT_CSV} ({len(seg)} image-mask pairs)") # ---- audit + EDA ----------------------------------------------------- audit = { "seed": SEED, "dataset_root": str(root), "n_images_indexed": int(len(df)), "n_images_readable": int(df["readable"].sum()), "n_images_with_mask": int(df["has_mask"].sum()), "n_missing_masks": len(issues.get("missing_masks", [])), "n_unreadable_images": len(issues.get("unreadable_images", [])), "n_unreadable_masks": len(issues.get("unreadable_masks", [])), "n_empty_masks": len(issues.get("empty_masks", [])), "n_non_binary_masks": len(issues.get("non_binary_masks", [])), "n_exact_duplicate_groups": len(issues.get("duplicate_images_exact", [])), "n_perceptual_duplicate_groups": len(issues.get("duplicate_images_perceptual", [])), "n_image_mask_size_mismatch": len(issues.get("image_mask_size_mismatch", [])), "split_counts": counts, "details": issues, } eda = eda_plots(d) audit["eda"] = eda save_json(audit, METRICS_DIR / f"A_data_audit_{SEED_TAG}.json") sel = qualitative_examples(d, n=args.n_examples) sel.to_csv(METRICS_DIR / f"A_seed_selected_examples_{SEED_TAG}.csv", index=False) print("\n=== AUDIT SUMMARY ===") for k in ("n_images_indexed", "n_images_with_mask", "n_missing_masks", "n_unreadable_images", "n_unreadable_masks", "n_empty_masks", "n_exact_duplicate_groups", "n_perceptual_duplicate_groups", "n_image_mask_size_mismatch"): print(f" {k:34s}: {audit[k]}") print(f" modal image size : {issues.get('modal_image_size')}") print(f" images at a non-modal size : {issues.get('n_unusual_image_sizes')}") print(f" imbalance ratio (max/min class) : {eda.get('imbalance_ratio_max_over_min'):.2f}") print("\nDone. Split files are frozen -- do not regenerate them after training starts.") if __name__ == "__main__": main()