440 lines
11 KiB
Python
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,
|
|
) |