product-detection/balance_and_augment.py

454 lines
18 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
balance_and_augment.py — Dataset Balancer & Augmentor
=======================================================
Analyses your crops_dataset/ (or data/) for class imbalance,
then generates augmented images until every class hits the
same target count.
Strategy
--------
1. Count images per class across train split
2. Find the target count (--target_count or auto = max class × scale)
3. For each under-represented class, generate augmented copies
until it reaches the target
4. Augmented images are written beside the originals — originals
are never touched
Augmentation pipeline (albumentations)
---------------------------------------
Light (always applied):
HorizontalFlip, slight rotation, brightness/contrast jitter,
hue/saturation shift
Medium (random subset):
Perspective, GridDistortion, GaussianBlur, Sharpen,
CoarseDropout (occlusion), JPEG compression noise
Heavy (low probability):
RandomSunFlare, RandomShadow, RandomFog, ChannelShuffle
Usage
-----
python balance_and_augment.py \
--data_dir crops_dataset \
--split train \
--target_count 300 \
--max_scale 3.0 \
--workers 8 \
--dry_run
Arguments
---------
--data_dir Root of the split dataset (contains train/ val/ test/)
--split Which split to augment (default: train)
--target_count Explicit target per class. Omit to use auto.
--max_scale Auto target = max_class_count × max_scale (default 2.0)
Capped so dominant classes are not inflated further.
--min_count Skip augmentation for classes already above this fraction
of the target (default 0.95 — within 5% is fine)
--workers Parallel workers (default: cpu_count - 1)
--dry_run Print the plan without writing any files
--val_split Also balance val split proportionally (0 = skip, default 0)
--seed Random seed for reproducibility (default 42)
--suffix Suffix added to augmented file stems (default: _aug)
--quality JPEG save quality for augmented images (default: 92)
"""
import argparse
import math
import multiprocessing
import random
import shutil
from collections import Counter
from pathlib import Path
import cv2
import numpy as np
try:
import albumentations as A
except ImportError:
raise ImportError("Run: pip install albumentations")
# ─────────────────────────── args ────────────────────────────────────────────
def parse_args():
p = argparse.ArgumentParser(
description="Balance and augment a classification dataset",
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
p.add_argument("--data_dir", default="crops_dataset")
p.add_argument("--split", default="train")
p.add_argument("--target_count", type=int, default=None,
help="Target images per class. Auto if omitted.")
p.add_argument("--max_scale", type=float, default=2.0,
help="Auto target = largest_class × max_scale")
p.add_argument("--min_count", type=float, default=0.95,
help="Classes already at ≥ this fraction of target are skipped")
p.add_argument("--workers", type=int,
default=max(1, multiprocessing.cpu_count() - 1))
p.add_argument("--dry_run", action="store_true",
help="Print plan only, write nothing")
p.add_argument("--val_split", type=float, default=0.0,
help="Fraction of augmented images to copy to val/ as well")
p.add_argument("--seed", type=int, default=42)
p.add_argument("--suffix", default="_aug",
help="Suffix added to augmented image filenames")
p.add_argument("--quality", type=int, default=92)
return p.parse_args()
# ─────────────────────────── augmentation pipeline ───────────────────────────
def build_pipeline(img_size: int) -> A.Compose:
"""
Three-tier augmentation pipeline tuned for product crops.
Each tier is applied with a probability so every generated
image is unique but still realistic.
"""
return A.Compose([
# ── Tier 1 — always active ─────────────────────────────────────────
A.HorizontalFlip(p=0.5),
A.ShiftScaleRotate(
shift_limit=0.06,
scale_limit=0.12,
rotate_limit=15,
border_mode=cv2.BORDER_REFLECT_101,
p=0.8,
),
A.RandomBrightnessContrast(
brightness_limit=0.30,
contrast_limit=0.30,
p=0.85,
),
A.HueSaturationValue(
hue_shift_limit=14,
sat_shift_limit=25,
val_shift_limit=20,
p=0.75,
),
# ── Tier 2 — medium augmentations (random subset) ─────────────────
A.OneOf([
A.Perspective(scale=(0.04, 0.10), p=1.0),
A.GridDistortion(num_steps=4, distort_limit=0.25, p=1.0),
A.ElasticTransform(alpha=60, sigma=8, p=1.0),
], p=0.40),
A.OneOf([
A.GaussianBlur(blur_limit=(3, 5), p=1.0),
A.MotionBlur(blur_limit=5, p=1.0),
A.Sharpen(alpha=(0.1, 0.4), p=1.0),
], p=0.35),
A.OneOf([
A.GaussNoise(var_limit=(5, 30), p=1.0),
A.ISONoise(color_shift=(0.01, 0.05), intensity=(0.05, 0.20), p=1.0),
], p=0.30),
# Occlusion: simulate a product partially hidden by another
A.CoarseDropout(
max_holes=4,
max_height=int(img_size * 0.20),
max_width=int(img_size * 0.20),
min_holes=1,
fill_value=0,
p=0.25,
),
# JPEG compression artefacts (common with phone cameras)
A.ImageCompression(quality_lower=55, quality_upper=90, p=0.25),
# ── Tier 3 — lighting / environment (low probability) ─────────────
A.OneOf([
A.RandomShadow(shadow_roi=(0, 0, 1, 1),
num_shadows_lower=1, num_shadows_upper=2,
shadow_dimension=4, p=1.0),
A.RandomSunFlare(flare_roi=(0, 0, 1, 0.5),
angle_lower=0.5, num_flare_circles_lower=2,
num_flare_circles_upper=6,
src_radius=100, p=1.0),
], p=0.15),
A.RandomFog(fog_coef_lower=0.05, fog_coef_upper=0.20,
alpha_coef=0.06, p=0.10),
# Colour channel shuffle (catches model over-reliance on colour)
A.ChannelShuffle(p=0.08),
# Final mild colour normalisation drift
A.RGBShift(r_shift_limit=12, g_shift_limit=12, b_shift_limit=12, p=0.30),
])
# ─────────────────────────── worker ──────────────────────────────────────────
def _augment_worker(task: dict) -> dict:
"""
Subprocess worker — generates `n_needed` augmented copies of one class.
task = {
"class_name": str,
"src_paths": [Path, ...],
"dst_dir": Path,
"n_needed": int,
"suffix": str,
"quality": int,
"seed": int,
"img_size": int,
}
Returns {"class_name": str, "generated": int, "errors": [str]}
"""
rng = random.Random(task["seed"])
np.random.seed(task["seed"] % (2**32))
pipeline = build_pipeline(task["img_size"])
src_paths = task["src_paths"]
dst_dir = Path(task["dst_dir"])
suffix = task["suffix"]
quality = task["quality"]
n_needed = task["n_needed"]
generated = 0
errors = []
# Cycle through source images, applying a different random transform each time
for i in range(n_needed):
src_path = src_paths[i % len(src_paths)]
img = cv2.imread(str(src_path))
if img is None:
errors.append(f"Cannot read {src_path}")
continue
img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
try:
result = pipeline(image=img_rgb)
aug_rgb = result["image"]
aug_bgr = cv2.cvtColor(aug_rgb, cv2.COLOR_RGB2BGR)
except Exception as e:
errors.append(f"Augmentation failed for {src_path}: e")
continue
# Build a unique output name: original_stem + suffix + counter
stem = Path(src_path).stem
out_name = f"{stem}{suffix}{i:05d}.jpg"
out_path = dst_dir / out_name
# Avoid overwriting an existing file
if out_path.exists():
out_name = f"{stem}{suffix}{i:05d}_{rng.randint(0,9999):04d}.jpg"
out_path = dst_dir / out_name
cv2.imwrite(
str(out_path), aug_bgr,
[cv2.IMWRITE_JPEG_QUALITY, quality],
)
generated += 1
return {
"class_name": task["class_name"],
"generated": generated,
"errors": errors,
}
# ─────────────────────────── helpers ─────────────────────────────────────────
IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".bmp", ".webp", ".tiff"}
def get_images(cls_dir: Path) -> list[Path]:
return [p for p in cls_dir.iterdir() if p.suffix.lower() in IMAGE_EXTS]
def detect_img_size(cls_dir: Path) -> int:
"""Sample the first image to get its smaller dimension (proxy for crop size)."""
for p in cls_dir.iterdir():
if p.suffix.lower() in IMAGE_EXTS:
img = cv2.imread(str(p))
if img is not None:
return min(img.shape[:2])
return 112 # fallback
def print_table(rows: list[dict], target: int):
"""Print a nicely formatted class statistics table."""
col = max(len(r["class"]) for r in rows) + 2
header = f" {'Class':<{col}} {'Current':>9} {'Target':>9} {'To add':>9} Status"
sep = "" * len(header)
print(sep)
print(header)
print(sep)
for r in rows:
need = max(0, target - r["current"])
status = "✓ ok" if need == 0 else f"+ {need:,}"
bar = "" * min(20, int(20 * r["current"] / target))
print(f" {r['class']:<{col}} {r['current']:>9,} {target:>9,} {need:>9,} {status} {bar}")
print(sep)
total_now = sum(r["current"] for r in rows)
total_add = sum(max(0, target - r["current"]) for r in rows)
print(f" {'TOTAL':<{col}} {total_now:>9,} {total_add:>9,}")
print(sep)
# ─────────────────────────── main ────────────────────────────────────────────
def main():
args = parse_args()
random.seed(args.seed)
split_dir = Path(args.data_dir) / args.split
if not split_dir.exists():
raise FileNotFoundError(f"Split directory not found: {split_dir}")
# ── Scan class directories ────────────────────────────────────────────────
class_dirs = sorted(
d for d in split_dir.iterdir()
if d.is_dir() and not d.name.startswith("_")
)
if not class_dirs:
raise RuntimeError(f"No class folders found in {split_dir}")
counts = {d.name: len(get_images(d)) for d in class_dirs}
if not any(counts.values()):
raise RuntimeError("All class folders appear empty.")
max_count = max(counts.values())
min_count = min(counts.values())
# ── Determine target ──────────────────────────────────────────────────────
if args.target_count:
target = args.target_count
print(f"\n[plan] Explicit target: {target:,} images per class")
else:
target = min(int(max_count * args.max_scale), max_count * 3)
print(f"\n[plan] Auto target: {max_count:,} (max) × {args.max_scale} = {target:,} per class")
# ── Print plan table ──────────────────────────────────────────────────────
rows = [{"class": name, "current": cnt} for name, cnt in sorted(counts.items())]
print()
print_table(rows, target)
needs_aug = [
d for d in class_dirs
if counts[d.name] < target * args.min_count
]
if not needs_aug:
print("\n✓ All classes already meet the target — nothing to do.")
return
print(f"\n Classes to augment: {len(needs_aug)}/{len(class_dirs)}")
print(f" Imbalance ratio : {max_count/max(min_count,1):.1f}x → 1.0x after augmentation")
if args.dry_run:
print("\n [dry-run] No files written. Remove --dry_run to apply.\n")
return
# ── Sample image size (for CoarseDropout sizing) ──────────────────────────
img_size = detect_img_size(class_dirs[0])
# ── Build worker tasks ────────────────────────────────────────────────────
tasks = []
for cls_dir in needs_aug:
name = cls_dir.name
src_imgs = get_images(cls_dir)
n_needed = target - len(src_imgs)
if n_needed <= 0:
continue
tasks.append({
"class_name": name,
"src_paths": [str(p) for p in src_imgs],
"dst_dir": str(cls_dir),
"n_needed": n_needed,
"suffix": args.suffix,
"quality": args.quality,
"seed": args.seed + abs(hash(name)) % 10000,
"img_size": img_size,
})
total_to_generate = sum(t["n_needed"] for t in tasks)
print(f"\n[augment] Generating {total_to_generate:,} images across "
f"{len(tasks)} classes using {args.workers} worker(s)…\n")
# ── Run (parallel per class) ──────────────────────────────────────────────
results = []
if args.workers > 1:
with multiprocessing.Pool(processes=args.workers) as pool:
for i, res in enumerate(pool.imap_unordered(_augment_worker, tasks), 1):
results.append(res)
done = sum(r["generated"] for r in results)
pct = done / total_to_generate * 100
print(f" [{i}/{len(tasks)}] {res['class_name']:<30} "
f"+{res['generated']:,} "
f"({pct:.0f}% total)", flush=True)
else:
for i, task in enumerate(tasks, 1):
res = _augment_worker(task)
results.append(res)
done = sum(r["generated"] for r in results)
pct = done / total_to_generate * 100
print(f" [{i}/{len(tasks)}] {res['class_name']:<30} "
f"+{res['generated']:,} "
f"({pct:.0f}% total)", flush=True)
# ── Collect errors ────────────────────────────────────────────────────────
all_errors = [(r["class_name"], e) for r in results for e in r["errors"]]
total_gen = sum(r["generated"] for r in results)
# ── Optional val propagation ──────────────────────────────────────────────
val_copied = 0
if args.val_split > 0:
val_dir = Path(args.data_dir) / "val"
if val_dir.exists():
print(f"\n[val] Copying {args.val_split:.0%} of augmented images to val/…")
for res in results:
cls_name = res["class_name"]
aug_files = sorted(
p for p in (split_dir / cls_name).iterdir()
if args.suffix in p.stem and p.suffix.lower() in IMAGE_EXTS
)
n_copy = max(1, int(len(aug_files) * args.val_split))
dst_cls = val_dir / cls_name
dst_cls.mkdir(exist_ok=True)
for src in aug_files[:n_copy]:
shutil.copy2(str(src), str(dst_cls / src.name))
val_copied += 1
else:
print(f"\n ⚠ val/ not found at {val_dir} — skipping val propagation")
# ── Final report ──────────────────────────────────────────────────────────
print("\n" + "" * 60)
print(" Augmentation complete!")
print(f" Images generated : {total_gen:,}")
if val_copied:
print(f" Copied to val/ : {val_copied:,}")
if all_errors:
print(f"\n{len(all_errors)} error(s):")
for cls, err in all_errors[:10]:
print(f" [{cls}] {err}")
if len(all_errors) > 10:
print(f" … and {len(all_errors)-10} more")
# ── Verify final counts ───────────────────────────────────────────────────
print("\n Final class counts:")
col = max(len(d.name) for d in class_dirs) + 2
print(f" {'Class':<{col}} {'Before':>8} {'After':>8}")
print(" " + "" * (col + 20))
for d in sorted(class_dirs, key=lambda x: x.name):
before = counts[d.name]
after = len(get_images(d))
delta = after - before
flag = f" +{delta:,}" if delta > 0 else ""
print(f" {d.name:<{col}} {before:>8,} {after:>8,}{flag}")
print("" * 60)
print(f"\n ➜ Next step: python 2_train_classifier.py --data_dir {args.data_dir}")
if __name__ == "__main__":
main()