Ishaank18's picture
Upload via upload_to_hf.py
d0518d9 verified
Raw History Blame Contribute Delete
24 kB
"""
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: <class>/images + <class>/masks (v3+),
or a flat <class>/ 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()