290 lines
9.9 KiB
Python
290 lines
9.9 KiB
Python
"""
|
|
STEP 2 — 2_augment.py
|
|
====================
|
|
Expands your small labeled dataset using augmentations tuned
|
|
specifically for retail shelf / packaged-goods detection.
|
|
|
|
~4 images → ~200 augmented images per class
|
|
(configurable via AUG_PER_IMAGE below)
|
|
|
|
Install:
|
|
pip install albumentations opencv-python tqdm pyyaml
|
|
|
|
Usage:
|
|
python augment.py
|
|
"""
|
|
|
|
import os
|
|
import cv2
|
|
import yaml
|
|
import random
|
|
import shutil
|
|
import numpy as np
|
|
|
|
from pathlib import Path
|
|
from tqdm import tqdm
|
|
|
|
import albumentations as A
|
|
|
|
# ── Configuration ────────────────────────────────────────────────────────────────
|
|
|
|
SRC_IMAGES = Path("dataset/images/train")
|
|
SRC_LABELS = Path("dataset/labels/train")
|
|
|
|
OUT_IMAGES = Path("dataset/images/train") # Augmented added in-place
|
|
OUT_LABELS = Path("dataset/labels/train")
|
|
|
|
VAL_IMAGES = Path("dataset/images/val")
|
|
VAL_LABELS = Path("dataset/labels/val")
|
|
|
|
AUG_PER_IMAGE = 50 # augmented copies per original image
|
|
VAL_SPLIT = 0.15 # fraction held out for validation
|
|
SEED = 42
|
|
|
|
IMG_SIZE = 640
|
|
|
|
# ── Augmentation pipeline ─────────────────────────────────────────────────────
|
|
# Chosen specifically for packaged-goods on retail shelves:
|
|
# • Perspective/affine → products photographed at angles
|
|
# • Brightness/contrast → fluorescent vs natural shelf lighting
|
|
# • Hue/saturation → the Famico red looks different under warm vs cool light
|
|
# • Blur → camera shake, shallow depth of field
|
|
# • Noise → phone camera noise
|
|
# • Shadow → shelf edge shadows across packages
|
|
# • Cutout → partial occlusion by other products
|
|
# • CLAHE → over/under-exposed shots
|
|
|
|
def build_augmentation_pipeline() -> A.Compose:
|
|
bbox_params = A.BboxParams(
|
|
format = "yolo" , # cx, cy, w, h normalised 0-1
|
|
label_fields = ["class_ids"],
|
|
min_area = 0.001 , # drop boxes that shrink below 0.1% of image
|
|
min_visibility = 0.03 , # drop boxes that become >70% occluded
|
|
)
|
|
|
|
return A.Compose([
|
|
# ── Geometry ────────────────────────────────────────────────────────
|
|
A.HorizontalFlip(p=0.5),
|
|
|
|
A.ShiftScaleRotate(
|
|
shift_limit = 0.08,
|
|
scale_limit = 0.25,
|
|
rotate_limit = 12,
|
|
border_mode = cv2.BORDER_CONSTANT,
|
|
value = 0,
|
|
p = 0.8
|
|
),
|
|
|
|
A.Perspective(
|
|
scale = (0.03, 0.10),
|
|
keep_size = True,
|
|
p = 0.5,
|
|
),
|
|
|
|
A.RandomResizedCrop(
|
|
size = (IMG_SIZE, IMG_SIZE),
|
|
scale = (0.70, 1.00),
|
|
ratio = (0.75, 1.33),
|
|
p = 0.4,
|
|
),
|
|
|
|
A.RandomBrightnessContrast(
|
|
brightness_limit = 0.35,
|
|
constrast_limit = 0.35,
|
|
p = 0.8,
|
|
),
|
|
|
|
A.HueSaturationValue(
|
|
hue_shift_limit = 12,
|
|
sat_shift_limit = 30,
|
|
val_shit_limit = 25,
|
|
p = 0.7,
|
|
),
|
|
|
|
A.CLAHE(
|
|
clip_limit = 4.0,
|
|
tile_grid_size = (8, 8),
|
|
p = 0.3,
|
|
),
|
|
|
|
A.RGBShift(
|
|
r_shift_limit = 15,
|
|
g_shift_limit = 10,
|
|
b_shift_limit = 10,
|
|
p = 0.4,
|
|
),
|
|
|
|
# ── Blur / Noise ─────────────────────────────────────────────────────────────────
|
|
A.OneOf([
|
|
A.GaussianBlur(blur_limit=(3, 7)),
|
|
A.MotionBlur (blur_limit=7),
|
|
A.MedianBlur (blur_limit=5),
|
|
], p=0.5),
|
|
|
|
A.GaussNoise(
|
|
std_range = (0.01, 0.05),
|
|
p = 0.4,
|
|
),
|
|
|
|
A.ImageCompression(
|
|
quality_range = (60, 95),
|
|
p = 0.4,
|
|
),
|
|
|
|
# ── Occlusion Simulation ──────────────────────────────────────────────────────────
|
|
A.CoarseDropout(
|
|
num_holes_range = (1, 4),
|
|
hole_height_range = (0.04, 0.12),
|
|
hole_width_range = (0.04, 0.12),
|
|
fill = 0,
|
|
p = 0.35,
|
|
),
|
|
|
|
# ── Final Resize to YOLO input size ───────────────────────────────────────────────
|
|
A.LongestMaxSize(max_size=IMG_SIZE),
|
|
A.PadIfNeeded(
|
|
min_height = IMG_SIZE,
|
|
min_width = IMG_SIZE,
|
|
border_mode = cv2.BORDER_CONSTANT,
|
|
value = 112, # Standard YOLO gray padding.
|
|
),
|
|
], bbox_params=bbox_params)
|
|
|
|
|
|
# ── YOLO label helpers ────────────────────────────────────────────────────────
|
|
|
|
def read_yolo_label(label_path: Path):
|
|
""" Retruns (class_ids: list[int], bboxes: list[tuple])"""
|
|
class_ids, bboxes = [], []
|
|
if not label_path.exists():
|
|
return class_ids, bboxes
|
|
|
|
with open(label_path) as f:
|
|
for line in f:
|
|
parts = line.strip().split()
|
|
if len(parts) != 5:
|
|
continue
|
|
|
|
cid = int(parts[0])
|
|
bbox = tuple(float(x) for x in parts[1:])
|
|
|
|
class_ids.append(cid)
|
|
bboxes .append(bbox)
|
|
|
|
return class_ids, bboxes
|
|
|
|
|
|
def write_yolo_label(label_path: Path, class_ids, bboxes):
|
|
with open(label_path, "w") as f:
|
|
for cid, (cx, cy, w, h) in zip(class_ids, bboxes):
|
|
f.write(f"{cid} {cx:6f} {cy:6f} {w:.6f} {h:.6f}\n")
|
|
|
|
def clamp_bbox(bbox):
|
|
cx, cy, w, h = bbox
|
|
cx = max(0.0, min(1.0, cx))
|
|
cy = max(0.0, min(1.0, cy))
|
|
|
|
w = max(0.001, min(1.0 - cx + w/2, w))
|
|
h = max(0.001, min(1.0 - cy + h/2, h))
|
|
|
|
return (cx, cy, w, h)
|
|
|
|
# ── Train / val split ─────────────────────────────────────────────────────────
|
|
|
|
def split_val(image_paths, val_ratio=VAL_SPLIT):
|
|
random.seed(SEED)
|
|
shuffled = list(image_paths)
|
|
random.shuffle(shuffled)
|
|
|
|
n_val = max(1, int(len(shuffled) * val_ratio))
|
|
return shuffled[n_val:], shuffled[:n_val] # train, val
|
|
|
|
def copy_to_val(image_path: Path):
|
|
stem = image_path.stem
|
|
label_path = SRC_LABELS / f"{stem}.txt"
|
|
|
|
VAL_IMAGES.mkdir(parents=True, exist_ok=True)
|
|
VAL_LABELS.mkdir(parents=True, exist_ok=True)
|
|
|
|
shutil.copy(image_path, VAL_IMAGES / image_path.name)
|
|
if label_path.exists():
|
|
shutil.copy(label_path, VAL_LABELS / f"{stem}.txt")
|
|
|
|
# ── Train / val split ─────────────────────────────────────────────────────────
|
|
def main():
|
|
random.seed(SEED)
|
|
np.random.seed(SEED)
|
|
|
|
exts = { ".jpg", ".jpeg", ".png", ".bmp", ".webp"}
|
|
originals = sorted([
|
|
p for p in SRC_IMAGES.iterdir()
|
|
if p.suffix.lower() in exts and not p.stem.startswith("aug_")
|
|
])
|
|
|
|
if not originals:
|
|
print(f"[ERROR] No images found in {SRC_IMAGES}")
|
|
print(" Put your labeled images there first, then run this script.")
|
|
return
|
|
|
|
print(f"Found {len(originals)} original images.")
|
|
|
|
# Split before augmenting - val set contains ONLY original photos
|
|
train_paths, val_paths = split_val(originals, VAL_SPLIT)
|
|
print(f"Validation set: {len(val_paths)} originals (never augmented)")
|
|
print(f"Training set: {len(train_paths)} originals -> " f"{len(train_paths) * AUG_PER_IMAGE} augmented\n")
|
|
|
|
for vp in val_paths:
|
|
copy_to_val(vp)
|
|
|
|
transform = build_augmentation_pipeline()
|
|
|
|
total_written = 0
|
|
for img_path in tqdm(train_paths, desc="Augmenting"):
|
|
image = cv2.imread(str(img_path))
|
|
if image is None:
|
|
print(f" [WANR] Cannot read {img_path}, skipping.")
|
|
continue
|
|
|
|
image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
|
|
class_ids, bboxes = read_yolo_label(
|
|
SRC_LABELS / f"{img_path.stem}.txt"
|
|
)
|
|
if not bboxes:
|
|
print(f" [WARN] No Labels for {img_path.name}, skipping.")
|
|
continue
|
|
|
|
for i in range(AUG_PER_IMAGE):
|
|
try:
|
|
result = transform(
|
|
image = image,
|
|
bboxes = bboxes,
|
|
class_ids = class_ids,
|
|
)
|
|
except Exception as e:
|
|
print(f" [WARN] Augmentation failed for {img_path.name} #{i}: {e}")
|
|
continue
|
|
|
|
aug_bboxes = [clamp_bbox(b) for b in result["bboxes"]]
|
|
aug_cids = result["class_ids"]
|
|
|
|
if not aug_bboxes:
|
|
continue # All Boxes fell out of frame.
|
|
|
|
aug_name = f"aug_{img_path.stem}_{i:04d}"
|
|
out_img = OUT_IMAGES / f"{aug_name}.jpg"
|
|
out_label = OUT_LABELS / f"{aug_name}.txt"
|
|
|
|
aug_img = cv2.cvtColor(result["image"], cv2.COLOR_RGB2BGR)
|
|
cv2.imwrite(str(out_img), aug_img, [cv2.IMWRITE_JPEG_QUALITY, 92])
|
|
write_yolo_label(out_label, aug_cids, aug_bboxes)
|
|
total_written += 1
|
|
|
|
print(f"\nDone - {total_written} augmented images written.")
|
|
print(f" Train: {OUT_IMAGES} Val: {VAL_IMAGES}")
|
|
print("\nRun next: python train.py")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|
|
|