agrobot_base/Python/OAK/datasets/oak-fcc-3/_5_augmentation.py

440 lines
11 KiB
Python

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
_5_augmentation_grouped_oak.py
Augmenta previews e máscaras por grupo na estrutura nova OAK-FCC-3.
Entrada:
dataset/originals/group/<grupo>/
previews/
masks/
metas/
bins/
Saída:
dataset/augmented/group/<grupo>/
previews/
masks/
Uso:
python _5_augmentation_grouped_oak.py --copies 5
python _5_augmentation_grouped_oak.py ^
--copies 5 ^
--groups chao,chao_cana,chao_erva,chao_cana_erva
python _5_augmentation_grouped_oak.py ^
--copies 5 ^
--src-root dataset/originals/group ^
--dst-root dataset/augmented/group ^
--clear-dst
"""
import os
import json
import argparse
import random
import shutil
from pathlib import Path
import cv2
from PIL import Image
import albumentations as A
# ====================== Configurações ======================
with open("config.json", "r", encoding="utf-8") as f:
config = json.load(f)
MODELO = config.get("camera", ".")
DATASET_BASE = os.path.join(MODELO, "dataset")
DEFAULT_SRC_ROOT = os.path.join(DATASET_BASE, "originals", "group")
DEFAULT_DST_ROOT = os.path.join(DATASET_BASE, "augmented", "group")
IMG_EXTS = (".jpg", ".jpeg", ".png")
MSK_EXTS = (".png", ".jpg", ".jpeg")
# ====================== Augmentation ======================
train_tf = A.Compose(
[
A.HorizontalFlip(p=0.5),
A.ShiftScaleRotate(
shift_limit=0.01,
scale_limit=0.10,
rotate_limit=5,
border_mode=cv2.BORDER_REFLECT_101,
interpolation=cv2.INTER_LINEAR,
mask_interpolation=cv2.INTER_NEAREST,
p=0.30,
),
A.OneOf(
[
A.RandomBrightnessContrast(0.2, 0.2, p=1.0),
A.HueSaturationValue(
hue_shift_limit=5,
sat_shift_limit=20,
val_shift_limit=15,
p=1.0,
),
A.RandomGamma(gamma_limit=(90, 110), p=1.0),
],
p=0.70,
),
A.OneOf(
[
A.MotionBlur(blur_limit=3, p=1.0),
A.GaussianBlur(blur_limit=3, p=1.0),
],
p=0.20,
),
A.OneOf(
[
A.GaussNoise(var_limit=(5.0, 15.0), p=1.0),
A.ImageCompression(quality_lower=60, quality_upper=90, p=1.0),
],
p=0.20,
),
A.RandomShadow(p=0.08),
A.RandomSunFlare(p=0.06),
# Cuidado com ChannelShuffle em imagens agrícolas:
# pode bagunçar demais a semântica das cores.
# Mantive desligado por padrão.
# A.ChannelShuffle(p=0.03),
A.CoarseDropout(
max_holes=6,
max_height=16,
max_width=16,
p=0.08,
),
]
)
# ====================== Utilitários ======================
def garantir_dir(p):
os.makedirs(p, exist_ok=True)
def maybe_clear_dir(p: str):
if os.path.isdir(p):
shutil.rmtree(p)
garantir_dir(p)
def normalizar_base(stem: str) -> str:
sufixos = [
"_rgb", "_RGB", "_Rgb",
"_image", "_img", "_frame",
"_preview", "_previews",
"_mask", "_masks",
"_seg", "_SEG", "_segment", "_segmentacao", "_Segmentacao",
]
out = stem
mudou = True
while mudou:
mudou = False
for sfx in sufixos:
if out.endswith(sfx):
out = out[: -len(sfx)]
mudou = True
break
return out
def map_files_by_base(folder, exts):
by_base = {}
if not os.path.isdir(folder):
return by_base
prioridade = {
".png": 0,
".jpg": 1,
".jpeg": 2,
}
for fname in os.listdir(folder):
lower = fname.lower()
if not lower.endswith(exts):
continue
stem, ext = os.path.splitext(fname)
base = normalizar_base(stem)
cand = os.path.join(folder, fname)
if base not in by_base:
by_base[base] = cand
else:
cur_ext = os.path.splitext(by_base[base])[1].lower()
if prioridade.get(ext.lower(), 99) < prioridade.get(cur_ext, 99):
by_base[base] = cand
return by_base
def list_groups(src_root):
if not os.path.isdir(src_root):
return []
grupos = []
for name in sorted(os.listdir(src_root)):
gdir = os.path.join(src_root, name)
if not os.path.isdir(gdir):
continue
previews = os.path.join(gdir, "previews")
masks = os.path.join(gdir, "masks")
if os.path.isdir(previews) and os.path.isdir(masks):
grupos.append(name)
return grupos
def load_rgb(path):
im = cv2.imread(path, cv2.IMREAD_COLOR)
if im is None:
raise FileNotFoundError(path)
return cv2.cvtColor(im, cv2.COLOR_BGR2RGB)
def save_rgb(path, arr_rgb):
garantir_dir(os.path.dirname(path))
Image.fromarray(arr_rgb).save(path)
def ensure_aug_dirs(dst_root, group_name):
preview_out = os.path.join(dst_root, group_name, "previews")
mask_out = os.path.join(dst_root, group_name, "masks")
garantir_dir(preview_out)
garantir_dir(mask_out)
return preview_out, mask_out
def validate_mask_colors_before_after(mask_before, mask_after, name="mask"):
"""
Só diagnóstico leve: máscaras não devem ganhar milhares de cores.
Em segmentação por cor, interpolação errada cria sujeira cromática.
"""
before_colors = len(set(map(tuple, mask_before.reshape(-1, 3))))
after_colors = len(set(map(tuple, mask_after.reshape(-1, 3))))
if after_colors > max(before_colors * 3, 32):
print(
f"[WARN] {name}: muitas cores após aug. "
f"antes={before_colors}, depois={after_colors}. "
f"Confira se mask_interpolation está NEAREST."
)
# ====================== Augmentação ======================
def augment_pair(
preview_path,
mask_path,
preview_out_dir,
mask_out_dir,
copies,
):
base_preview, preview_ext = os.path.splitext(os.path.basename(preview_path))
_, mask_ext = os.path.splitext(os.path.basename(mask_path))
base = normalizar_base(base_preview)
img = load_rgb(preview_path)
mask = load_rgb(mask_path)
gen = 0
for i in range(copies):
aug = train_tf(image=img, mask=mask)
img_aug = aug["image"]
mask_aug = aug["mask"]
validate_mask_colors_before_after(mask, mask_aug, name=f"{base}_aug_{i:02d}")
new_base = f"{base}_aug_{i:02d}"
out_preview = os.path.join(preview_out_dir, f"{new_base}{preview_ext}")
out_mask = os.path.join(mask_out_dir, f"{new_base}{mask_ext}")
save_rgb(out_preview, img_aug)
save_rgb(out_mask, mask_aug)
gen += 1
return gen
# ====================== Processamento ======================
def process_group(
src_root,
dst_root,
group_name,
copies,
limit=None,
seed=42,
):
group_dir = os.path.join(src_root, group_name)
preview_dir = os.path.join(group_dir, "previews")
mask_dir = os.path.join(group_dir, "masks")
if not os.path.isdir(preview_dir) or not os.path.isdir(mask_dir):
print(f"[WARN] Grupo '{group_name}' inválido, sem previews/masks. Pulando.")
return 0
preview_map = map_files_by_base(preview_dir, IMG_EXTS)
mask_map = map_files_by_base(mask_dir, MSK_EXTS)
bases = sorted(set(preview_map.keys()) & set(mask_map.keys()))
if limit is not None and limit > 0 and limit < len(bases):
rng = random.Random(seed)
bases = sorted(rng.sample(bases, limit))
print(f"[INFO] [{group_name}] Limit aplicado: {limit} amostras originais selecionadas.")
preview_out_dir, mask_out_dir = ensure_aug_dirs(dst_root, group_name)
count = 0
sem_mask = 0
erros = 0
for base in bases:
preview_path = preview_map.get(base)
mask_path = mask_map.get(base)
if not preview_path or not mask_path:
sem_mask += 1
print(f"[WARN] [{group_name}] par incompleto para base={base}. Pulando.")
continue
try:
count += augment_pair(
preview_path=preview_path,
mask_path=mask_path,
preview_out_dir=preview_out_dir,
mask_out_dir=mask_out_dir,
copies=copies,
)
except Exception as e:
erros += 1
print(f"[ERRO] [{group_name}] {base}: {e}")
print(
f"[OK] Grupo '{group_name}'{count} pares gerados. "
f"originais={len(bases)} | sem_mask={sem_mask} | erros={erros}"
)
return count
def main(
copies=5,
groups_csv=None,
src_root=DEFAULT_SRC_ROOT,
dst_root=DEFAULT_DST_ROOT,
limit=None,
seed=42,
clear_dst=False,
):
total = 0
print(f"[INFO] MODELO={MODELO}")
print(f"[INFO] SRC_ROOT={src_root}")
print(f"[INFO] DST_ROOT={dst_root}")
print(f"[INFO] copies={copies}")
print(f"[INFO] single_head=True")
if not os.path.isdir(src_root):
raise SystemExit(f"[ERRO] src-root não encontrado: {src_root}")
if clear_dst:
print(f"[INFO] Limpando destino: {dst_root}")
maybe_clear_dir(dst_root)
grupos = list_groups(src_root)
if groups_csv:
want = {g.strip() for g in groups_csv.split(",") if g.strip()}
grupos = [g for g in grupos if g in want]
if not grupos:
print("[WARN] Nenhum grupo válido encontrado após filtro.")
if not grupos:
print("[WARN] Nenhum grupo encontrado.")
return
print(f"Grupos encontrados: {', '.join(grupos)}")
for g in grupos:
total += process_group(
src_root=src_root,
dst_root=dst_root,
group_name=g,
copies=copies,
limit=limit,
seed=seed,
)
print(f"\nAugmentation completed! Total: {total} pares gerados.")
if __name__ == "__main__":
ap = argparse.ArgumentParser(
description="Augmentação por grupos OAK-FCC-3 usando previews/masks, single head."
)
ap.add_argument("--copies", type=int, default=5, help="Número de cópias augmentadas por imagem.")
ap.add_argument("--groups", type=str, default=None, help="Lista de grupos separados por vírgula.")
ap.add_argument("--src-root", type=str, default=DEFAULT_SRC_ROOT, help="Raiz dos grupos originais.")
ap.add_argument("--dst-root", type=str, default=DEFAULT_DST_ROOT, help="Raiz dos grupos augmentados.")
ap.add_argument("--limit", type=int, default=None, help="Quantidade máxima de amostras originais por grupo para augmentar.")
ap.add_argument("--seed", type=int, default=42, help="Seed para seleção reproduzível quando usar --limit.")
ap.add_argument("--clear-dst", action="store_true", help="Apaga dst-root antes de gerar.")
args = ap.parse_args()
random.seed(args.seed)
main(
copies=args.copies,
groups_csv=args.groups,
src_root=args.src_root,
dst_root=args.dst_root,
limit=args.limit,
seed=args.seed,
clear_dst=args.clear_dst,
)