agrobot_base/Python/OAK/datasets/_5_augmentation.py

316 lines
11 KiB
Python
Raw Normal View History

2025-09-15 10:23:07 +00:00
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Augmenta imagens e máscaras *por grupo*.
Entrada (via config.json -> MODELO):
MODELO/dataset/original/group/<grupo>/images
MODELO/dataset/original/group/<grupo>/masks
Saída:
MODELO/dataset/augmented/group/<grupo>/images
MODELO/dataset/augmented/group/<grupo>/masks
Se "original/group" não existir, faz fallback para:
MODELO/dataset/original/{images,masks}
MODELO/dataset/augmented/{images,masks}
Transf. geométricas (aplicam a img e máscara) e fotométricas (apenas imagem).
Uso:
python _3_augmentation_grouped.py --copies 5
python _3_augmentation_grouped.py --copies 5 --groups chao,chao_erva,cana
"""
import os
import json
import cv2
from PIL import Image
import albumentations as A
import argparse
# ⚙️ Configurações
with open("config_oak.json", "r", encoding="utf-8") as f:
2025-09-15 10:23:07 +00:00
config = json.load(f)
MODELO = config.get("camera", ".")
USE_MASKS2 = config.get("dual_head", False)
2025-09-15 10:23:07 +00:00
# Pastas base
DATASET_BASE = os.path.join(MODELO, "dataset")
ORIG_GROUP_ROOT = os.path.join(DATASET_BASE, "original", "group")
AUG_GROUP_ROOT = os.path.join(DATASET_BASE, "augmented", "group")
# Fallback (modo antigo, sem grupos)
ORIG_OLD_IMG = os.path.join(DATASET_BASE, "original", "images")
ORIG_OLD_MSK = os.path.join(DATASET_BASE, "original", "masks")
ORIG_OLD_MSK2 = os.path.join(DATASET_BASE, "original", "masks2")
2025-09-15 10:23:07 +00:00
AUG_OLD_IMG = os.path.join(DATASET_BASE, "augmented", "images")
AUG_OLD_MSK = os.path.join(DATASET_BASE, "augmented", "masks")
AUG_OLD_MSK2 = os.path.join(DATASET_BASE, "augmented", "masks2")
2025-09-15 10:23:07 +00:00
# Extensões aceitas
IMG_EXTS = (".jpg", ".jpeg", ".png")
MSK_EXTS = (".png", ".jpg", ".jpeg") # manter prioridade PNG quando possível
MSK2_EXTS = (".png", ".jpg", ".jpeg")
2025-09-15 10:23:07 +00:00
def garantir_dir(p):
os.makedirs(p, exist_ok=True)
# Pipeline de augmentations
train_tf = A.Compose([
A.HorizontalFlip(p=0.5),
# Geométricas (aplicam em imagem e máscara)
A.ShiftScaleRotate(
shift_limit=0.01,
scale_limit=0.10,
rotate_limit=5,
border_mode=cv2.BORDER_REFLECT_101,
interpolation=cv2.INTER_LINEAR,
p=0.30
),
# Fotométricas (somente imagem)
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=50, quality_upper=85, p=1.0),
], p=0.20),
A.RandomShadow(p=0.10),
A.RandomSunFlare(p=0.10),
A.ChannelShuffle(p=0.05),
A.CoarseDropout(max_holes=6, max_height=16, max_width=16, p=0.10),
], additional_targets={
'mask2': 'mask'
})
2025-09-15 10:23:07 +00:00
def load_rgb(path):
# cv2 lê BGR → converte para RGB
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):
Image.fromarray(arr_rgb).save(path)
def list_groups(root):
"""Lista grupos válidos (que contêm subpastas images e masks)."""
if not os.path.isdir(root):
return []
grupos = []
for name in sorted(os.listdir(root)):
gdir = os.path.join(root, name)
if not os.path.isdir(gdir):
continue
if os.path.isdir(os.path.join(gdir, "images")) and os.path.isdir(os.path.join(gdir, "masks")):
grupos.append(name)
return grupos
def map_masks_by_base(msk_dir):
"""Mapeia máscaras por base (prioriza .png)."""
by_base = {}
if not os.path.isdir(msk_dir):
return by_base
for fname in os.listdir(msk_dir):
f_lower = fname.lower()
if not f_lower.endswith(MSK_EXTS):
continue
base, ext = os.path.splitext(fname)
cand = os.path.join(msk_dir, fname)
if base not in by_base:
by_base[base] = cand
else:
# mantém .png se disponível
cur_ext = os.path.splitext(by_base[base])[1].lower()
if cur_ext != ".png" and ext.lower() == ".png":
by_base[base] = cand
return by_base
def map_masks2_by_base(msk2_dir):
"""Mapeia máscaras2 por base (prioriza .png)."""
by_base = {}
if not os.path.isdir(msk2_dir):
return by_base
for fname in os.listdir(msk2_dir):
f_lower = fname.lower()
if not f_lower.endswith(MSK2_EXTS):
continue
base, ext = os.path.splitext(fname)
cand = os.path.join(msk2_dir, fname)
if base not in by_base:
by_base[base] = cand
else:
cur_ext = os.path.splitext(by_base[base])[1].lower()
if cur_ext != ".png" and ext.lower() == ".png":
by_base[base] = cand
return by_base
def ensure_aug_dirs(group_name=None, use_masks2=False):
"""Cria diretórios de saída para o grupo ou modo antigo. Se use_masks2, cria masks2."""
2025-09-15 10:23:07 +00:00
if group_name:
img_out = os.path.join(AUG_GROUP_ROOT, group_name, "images")
msk_out = os.path.join(AUG_GROUP_ROOT, group_name, "masks")
msk2_out = os.path.join(AUG_GROUP_ROOT, group_name, "masks2") if use_masks2 else None
2025-09-15 10:23:07 +00:00
else:
img_out = AUG_OLD_IMG
msk_out = AUG_OLD_MSK
msk2_out = AUG_OLD_MSK2 if use_masks2 else None
2025-09-15 10:23:07 +00:00
garantir_dir(img_out)
garantir_dir(msk_out)
if use_masks2 and msk2_out:
garantir_dir(msk2_out)
return img_out, msk_out, msk2_out
2025-09-15 10:23:07 +00:00
def augment_pair(img_path, msk_path, img_out_dir, msk_out_dir, copies, msk2_path=None, msk2_out_dir=None):
2025-09-15 10:23:07 +00:00
base_img, img_ext = os.path.splitext(os.path.basename(img_path))
base_msk, msk_ext = os.path.splitext(os.path.basename(msk_path))
msk2_ext = os.path.splitext(os.path.basename(msk2_path))[1] if msk2_path else None
2025-09-15 10:23:07 +00:00
# padroniza pelo base da imagem
base = base_img
img = load_rgb(img_path)
msk = load_rgb(msk_path)
msk2 = load_rgb(msk2_path) if msk2_path else None
2025-09-15 10:23:07 +00:00
gen = 0
for i in range(copies):
if msk2 is not None and msk2_out_dir:
aug = train_tf(image=img, mask=msk, mask2=msk2)
else:
aug = train_tf(image=img, mask=msk)
2025-09-15 10:23:07 +00:00
img_aug = aug["image"]
msk_aug = aug["mask"]
out_img = os.path.join(img_out_dir, f"{base}_aug_{i:02d}{img_ext}")
out_msk = os.path.join(msk_out_dir, f"{base}_aug_{i:02d}{msk_ext}")
save_rgb(out_img, img_aug)
save_rgb(out_msk, msk_aug)
if msk2 is not None and msk2_out_dir:
msk2_aug = aug["mask2"]
out_msk2 = os.path.join(msk2_out_dir, f"{base}_aug_{i:02d}{msk2_ext}")
save_rgb(out_msk2, msk2_aug)
2025-09-15 10:23:07 +00:00
gen += 1
return gen
def process_group(group_name, copies):
"""Processa um grupo único (images/masks dentro de ORIG_GROUP_ROOT/<group_name>/)."""
img_dir = os.path.join(ORIG_GROUP_ROOT, group_name, "images")
msk_dir = os.path.join(ORIG_GROUP_ROOT, group_name, "masks")
msk2_dir = os.path.join(ORIG_GROUP_ROOT, group_name, "masks2")
2025-09-15 10:23:07 +00:00
if not (os.path.isdir(img_dir) and os.path.isdir(msk_dir)):
print(f"[WARN] Grupo '{group_name}' inválido (sem images/masks). Pulando.")
return 0
imgs = [f for f in os.listdir(img_dir) if os.path.splitext(f.lower())[1] in IMG_EXTS]
msk_map = map_masks_by_base(msk_dir)
use_masks2 = USE_MASKS2 and os.path.isdir(msk2_dir)
msk2_map = map_masks2_by_base(msk2_dir) if use_masks2 else {}
img_out_dir, msk_out_dir, msk2_out_dir = ensure_aug_dirs(group_name, use_masks2=use_masks2)
2025-09-15 10:23:07 +00:00
count = 0
for img_file in sorted(imgs):
base, _ = os.path.splitext(img_file)
msk_file = msk_map.get(base)
if not msk_file:
print(f"[WARN] [{group_name}] Máscara não encontrada para {img_file}, pulando.")
continue
msk2_file = msk2_map.get(base) if use_masks2 else None
if use_masks2 and not msk2_file:
print(f"[WARN] [{group_name}] mask2 não encontrada para {img_file}, gerando só img+mask.")
2025-09-15 10:23:07 +00:00
try:
count += augment_pair(
os.path.join(img_dir, img_file),
msk_file,
img_out_dir,
msk_out_dir,
copies=copies,
msk2_path=msk2_file,
msk2_out_dir=msk2_out_dir
2025-09-15 10:23:07 +00:00
)
except Exception as e:
print(f"[ERRO] [{group_name}] {img_file}: {e}")
print(f"[OK] Grupo '{group_name}'{count} pares gerados.")
return count
def process_legacy(copies):
"""Fallback: modo sem grupos (original/images e original/masks)."""
if not (os.path.isdir(ORIG_OLD_IMG) and os.path.isdir(ORIG_OLD_MSK)):
print("[WARN] Modo legacy não encontrado. Nada a fazer.")
return 0
imgs = [f for f in os.listdir(ORIG_OLD_IMG) if os.path.splitext(f.lower())[1] in IMG_EXTS]
msk_map = map_masks_by_base(ORIG_OLD_MSK)
use_masks2 = USE_MASKS2 and os.path.isdir(ORIG_OLD_MSK2)
msk2_map = map_masks2_by_base(ORIG_OLD_MSK2) if use_masks2 else {}
img_out_dir, msk_out_dir, msk2_out_dir = ensure_aug_dirs(group_name=None, use_masks2=use_masks2)
2025-09-15 10:23:07 +00:00
count = 0
for img_file in sorted(imgs):
base, _ = os.path.splitext(img_file)
msk_file = msk_map.get(base)
if not msk_file:
print(f"[WARN] (legacy) Máscara não encontrada para {img_file}, pulando.")
continue
msk2_file = msk2_map.get(base) if use_masks2 else None
if use_masks2 and not msk2_file:
print(f"[WARN] (legacy) mask2 não encontrada para {img_file}, gerando só img+mask.")
2025-09-15 10:23:07 +00:00
try:
count += augment_pair(
os.path.join(ORIG_OLD_IMG, img_file),
msk_file,
img_out_dir,
msk_out_dir,
copies=copies,
msk2_path=msk2_file,
msk2_out_dir=msk2_out_dir
2025-09-15 10:23:07 +00:00
)
except Exception as e:
print(f"[ERRO] (legacy) {img_file}: {e}")
print(f"[OK] Legacy → {count} pares gerados.")
return count
def main(copies=5, groups_csv=None):
total = 0
if os.path.isdir(ORIG_GROUP_ROOT):
grupos = list_groups(ORIG_GROUP_ROOT)
if groups_csv:
# filtra pelos grupos desejados
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 em original/group. Tentando modo legacy...")
total += process_legacy(copies)
else:
print(f"Grupos encontrados: {', '.join(grupos)}")
for g in grupos:
total += process_group(g, copies)
else:
# sem estrutura de grupos
total += process_legacy(copies)
print(f"\nAugmentation completed! Total: {total} pares gerados.")
if __name__ == "__main__":
ap = argparse.ArgumentParser(description="Augmentação por grupos (images/masks)")
ap.add_argument("--copies", type=int, default=5, help="Número de cópias augmentadas por imagem (default=5).")
ap.add_argument("--groups", type=str, default=None, help="Lista de grupos separados por vírgula (ex: chao,erva_cana).")
args = ap.parse_args()
main(copies=args.copies, groups_csv=args.groups)