agrobot_base/Python/OAK/datasets/_6_normalize.py

239 lines
9.5 KiB
Python

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Normaliza/redimensiona imagens e máscaras mantendo a ESTRUTURA POR GRUPO.
Entradas (via config.json -> MODELO, RESOLUCAO):
- MODELO/dataset/original/group/<grupo>/{images,masks}
- MODELO/dataset/augmented/group/<grupo>/{images,masks}
Saídas (por resolução):
- MODELO/dataset/<WxH>/group/<grupo>/{images,masks}
Fallback (modo legado, se não houver "group/"):
- original/{images,masks} e augmented/{images,masks} -> <WxH>/{images,masks}
Conversão de máscara:
- Lê máscara RGB e converte para IDs via utils.converter_mask_rgb_para_ids
- Ignora classe "ignore" conforme labelmap (usa índice 255 como padrão quando necessário)
"""
import argparse
import os
import json
import cv2
from typing import Dict, List, Tuple
from utils import carregar_labelmap_completo, converter_mask_rgb_para_ids
# ⚙️ Configurações
with open("config.json", "r", encoding="utf-8") as f:
config = json.load(f)
MODELO = config["camera"]
RESOLUCAO = tuple(config["resolucao"]) # [W, H] ou [width, height]
pasta_base = os.path.join(MODELO, "dataset")
labelmap_path = os.path.join(pasta_base, "labelmap.txt")
# Dimensões alvo (pode expandir para múltiplas se quiser)
RESOLUCOES = {
f"{RESOLUCAO[0]}x{RESOLUCAO[1]}": (RESOLUCAO[0], RESOLUCAO[1]),
}
# Fontes a processar
FONTES = ["original", "augmented"]
# Extensões aceitas
IMG_EXTS = (".jpg", ".jpeg", ".png")
MSK_EXTS = (".png", ".jpg", ".jpeg") # preferir .png
def infer_ignore_id(ignore_rgb, default_id=255):
"""
Tenta inferir o ID de ignore a partir do valor retornado por carregar_labelmap_completo.
- Se for [id] retorna id
- Se for (R,G,B) retorna default_id (tipicamente 255)
- Se for int, retorna direto
"""
if isinstance(ignore_rgb, (list, tuple)):
if len(ignore_rgb) == 1:
try:
return int(ignore_rgb[0])
except Exception:
return default_id
if len(ignore_rgb) == 3:
return default_id
if isinstance(ignore_rgb, int):
return ignore_rgb
return default_id
def garantir_dir(p):
os.makedirs(p, exist_ok=True)
def list_groups(root) -> List[str]:
"""Lista grupos válidos com 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: str) -> Dict[str, str]:
"""Retorna {base: caminho_mask}, priorizando .png quando houver múltiplas por base."""
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:
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 normalize_pair(caminho_rgb: str, caminho_mask: str, cor_para_id, ignore_id: int,
out_img_dir: str, out_msk_dir: str, dim: Tuple[int,int], prefix: str = ""):
"""Redimensiona e grava a imagem e a máscara (se houver)."""
img_rgb = cv2.imread(caminho_rgb)
if img_rgb is None:
print(f"[!] Erro ao ler imagem: {caminho_rgb}")
return False
# Nome de saída com prefixo para distinguir fonte (ex: original_, augmented_)
nome = os.path.basename(caminho_rgb)
if prefix:
nome_saida_img = f"{prefix}{nome}"
else:
nome_saida_img = nome
nome_saida_msk = nome_saida_img
for ext in (".jpg", ".jpeg", ".png"):
if nome_saida_msk.lower().endswith(ext):
nome_saida_msk = nome_saida_msk[: -len(ext)] + ".png"
break
# Redimensiona imagem
img_resized = cv2.resize(img_rgb, dim, interpolation=cv2.INTER_AREA)
garantir_dir(out_img_dir)
cv2.imwrite(os.path.join(out_img_dir, nome_saida_img), img_resized)
# Processa e redimensiona máscara (se existir)
if caminho_mask and os.path.isfile(caminho_mask):
msk_bgr = cv2.imread(caminho_mask, cv2.IMREAD_COLOR)
if msk_bgr is None:
print(f"[!] Erro ao ler máscara: {caminho_mask}")
else:
msk_rgb = cv2.cvtColor(msk_bgr, cv2.COLOR_BGR2RGB)
mask_ids = converter_mask_rgb_para_ids(msk_rgb, cor_para_id, ignore_id)
mask_resized = cv2.resize(mask_ids, dim, interpolation=cv2.INTER_NEAREST)
garantir_dir(out_msk_dir)
cv2.imwrite(os.path.join(out_msk_dir, nome_saida_msk), mask_resized)
return True
def process_group_root(fonte_root: str, fonte_nome: str, cor_para_id, ignore_id: int, groups_except: str = None):
"""Processa uma raiz do tipo .../<fonte>/group/ agrupando por cada subpasta de grupo."""
total = 0
grupos = list_groups(fonte_root)
if not grupos:
return 0
not_want = {g.strip() for g in groups_except.split(",") if g.strip()}
grupos_desconsiderar = [g for g in grupos if g in not_want]
for nome_res, dim in RESOLUCOES.items():
out_root = os.path.join(pasta_base, nome_res, "group")
for grupo in grupos:
if grupo in grupos_desconsiderar:
print(f"[WARN] Grupo desconsiderado nao sera processado: {grupo}")
continue
in_img_dir = os.path.join(fonte_root, grupo, "images")
in_msk_dir = os.path.join(fonte_root, grupo, "masks")
if not (os.path.isdir(in_img_dir) and os.path.isdir(in_msk_dir)):
print(f"[WARN] Grupo inválido (sem images/masks): {grupo}")
continue
out_img_dir = os.path.join(out_root, grupo, "images")
out_msk_dir = os.path.join(out_root, grupo, "masks")
msk_map = map_masks_by_base(in_msk_dir)
imgs = [f for f in os.listdir(in_img_dir) if os.path.splitext(f.lower())[1] in IMG_EXTS]
n = len(imgs)
for i, fname in enumerate(sorted(imgs), 1):
base, _ = os.path.splitext(fname)
caminho_rgb = os.path.join(in_img_dir, fname)
caminho_mask = msk_map.get(base)
ok = normalize_pair(
caminho_rgb, caminho_mask, cor_para_id, ignore_id,
out_img_dir, out_msk_dir, dim, prefix=f"{fonte_nome}_"
)
if ok:
total += 1
print(f"[{fonte_nome} | {grupo} | {nome_res}] {i}/{n}{fname}")
return total
def process_legacy_root(legacy_img: str, legacy_msk: str, fonte_nome: str, cor_para_id, ignore_id: int):
"""Processa estrutura legado (sem grupos)."""
if not (os.path.isdir(legacy_img) and os.path.isdir(legacy_msk)):
return 0
total = 0
for nome_res, dim in RESOLUCOES.items():
out_img_dir = os.path.join(pasta_base, nome_res, "images")
out_msk_dir = os.path.join(pasta_base, nome_res, "masks")
msk_map = map_masks_by_base(legacy_msk)
imgs = [f for f in os.listdir(legacy_img) if os.path.splitext(f.lower())[1] in IMG_EXTS]
n = len(imgs)
for i, fname in enumerate(sorted(imgs), 1):
base, _ = os.path.splitext(fname)
caminho_rgb = os.path.join(legacy_img, fname)
caminho_mask = msk_map.get(base)
ok = normalize_pair(
caminho_rgb, caminho_mask, cor_para_id, ignore_id,
out_img_dir, out_msk_dir, dim, prefix=f"{fonte_nome}_"
)
if ok:
total += 1
print(f"[{fonte_nome} | legacy | {nome_res}] {i}/{n}{fname}")
return total
def main(args):
# === Labelmap ===
# Espera tupla na ordem: (cor_para_id, colormap_rgb, id_para_nome, ignore_rgb)
cor_para_id, _colormap_rgb, _id_para_nome, ignore_rgb = carregar_labelmap_completo(labelmap_path)
ignore_id = infer_ignore_id(ignore_rgb, default_id=255)
total_geral = 0
# === ORIGINAL ===
orig_group_root = os.path.join(pasta_base, "original", "group")
if os.path.isdir(orig_group_root):
total_geral += process_group_root(orig_group_root, "original", cor_para_id, ignore_id, groups_except=args.groups_except)
else:
legacy_img = os.path.join(pasta_base, "original", "images")
legacy_msk = os.path.join(pasta_base, "original", "masks")
total_geral += process_legacy_root(legacy_img, legacy_msk, "original", cor_para_id, ignore_id)
# === AUGMENTED ===
aug_group_root = os.path.join(pasta_base, "augmented", "group")
if os.path.isdir(aug_group_root):
total_geral += process_group_root(aug_group_root, "augmented", cor_para_id, ignore_id, groups_except=args.groups_except)
else:
legacy_img = os.path.join(pasta_base, "augmented", "images")
legacy_msk = os.path.join(pasta_base, "augmented", "masks")
total_geral += process_legacy_root(legacy_img, legacy_msk, "augmented", cor_para_id, ignore_id)
print(f"\n✅ Concluído! Total normalizados: {total_geral}")
if __name__ == "__main__":
ap = argparse.ArgumentParser(description="Augmentação por grupos (images/masks)")
ap.add_argument("--groups-except", type=str, default="", help="Lista de grupos para nao usar, separados por vírgula (ex: chao,erva_cana).")
args = ap.parse_args()
main(args)