2025-09-15 10:23:07 +00:00
|
|
|
# -*- coding: utf-8 -*-
|
2025-08-07 18:23:53 +00:00
|
|
|
from PIL import Image
|
|
|
|
|
import glob
|
|
|
|
|
import os
|
|
|
|
|
import cv2
|
|
|
|
|
import numpy as np
|
|
|
|
|
import torch
|
|
|
|
|
from torch.utils.data import Dataset
|
|
|
|
|
|
|
|
|
|
from utils import carregar_labelmap_completo, compute_roi_indices, resize_keep_width
|
|
|
|
|
|
2025-09-15 10:23:07 +00:00
|
|
|
IMG_EXTS = (".jpg", ".jpeg", ".png")
|
|
|
|
|
MSK_EXTS = (".png", ".jpg", ".jpeg") # preferimos .png se existir
|
|
|
|
|
|
|
|
|
|
def _is_dir(p): return os.path.isdir(p)
|
|
|
|
|
def _is_file(p): return os.path.isfile(p)
|
|
|
|
|
|
|
|
|
|
def _list_groups(group_root):
|
|
|
|
|
if not _is_dir(group_root): return []
|
|
|
|
|
out = []
|
|
|
|
|
for g in sorted(os.listdir(group_root)):
|
|
|
|
|
gdir = os.path.join(group_root, g)
|
|
|
|
|
if not _is_dir(gdir):
|
|
|
|
|
continue
|
|
|
|
|
if _is_dir(os.path.join(gdir, "images")) and _is_dir(os.path.join(gdir, "masks")):
|
|
|
|
|
out.append(g)
|
|
|
|
|
return out
|
|
|
|
|
|
|
|
|
|
def _mask_for_base(msk_dir, base):
|
|
|
|
|
"""Encontra a máscara que casa com o base, priorizando .png."""
|
|
|
|
|
best = None
|
|
|
|
|
for ext in MSK_EXTS:
|
|
|
|
|
cand = os.path.join(msk_dir, base + ext)
|
|
|
|
|
if _is_file(cand):
|
|
|
|
|
if best is None: best = cand
|
|
|
|
|
# mantém .png se aparecer depois
|
|
|
|
|
if os.path.splitext(cand)[1].lower() == ".png":
|
|
|
|
|
return cand
|
|
|
|
|
return best
|
|
|
|
|
|
|
|
|
|
def _collect_pairs_legacy(root):
|
|
|
|
|
"""root/{images,masks}"""
|
|
|
|
|
img_dir = os.path.join(root, "images")
|
|
|
|
|
msk_dir = os.path.join(root, "masks")
|
|
|
|
|
imgs = []
|
|
|
|
|
msks = []
|
|
|
|
|
for p in sorted(glob.glob(os.path.join(img_dir, "*"))):
|
|
|
|
|
base, ext = os.path.splitext(os.path.basename(p))
|
|
|
|
|
if ext.lower() not in IMG_EXTS:
|
|
|
|
|
continue
|
|
|
|
|
m = _mask_for_base(msk_dir, base)
|
|
|
|
|
if m:
|
|
|
|
|
imgs.append(p)
|
|
|
|
|
msks.append(m)
|
|
|
|
|
return imgs, msks
|
|
|
|
|
|
|
|
|
|
def _collect_pairs_grouped(root):
|
|
|
|
|
"""Suporta:
|
|
|
|
|
- root/group/<g>/{images,masks}
|
|
|
|
|
- root/<g>/{images,masks} (quando root já é '.../group')
|
|
|
|
|
"""
|
|
|
|
|
# case A: root tem subpasta 'group'
|
|
|
|
|
group_root = os.path.join(root, "group")
|
|
|
|
|
if not _is_dir(group_root):
|
|
|
|
|
# case B: root JÁ É a pasta 'group'
|
|
|
|
|
group_root = root
|
|
|
|
|
|
|
|
|
|
groups = _list_groups(group_root)
|
|
|
|
|
imgs, msks = [], []
|
|
|
|
|
for g in groups:
|
|
|
|
|
img_dir = os.path.join(group_root, g, "images")
|
|
|
|
|
msk_dir = os.path.join(group_root, g, "masks")
|
|
|
|
|
for p in sorted(glob.glob(os.path.join(img_dir, "*"))):
|
|
|
|
|
base, ext = os.path.splitext(os.path.basename(p))
|
|
|
|
|
if ext.lower() not in IMG_EXTS:
|
|
|
|
|
continue
|
|
|
|
|
m = _mask_for_base(msk_dir, base)
|
|
|
|
|
if m:
|
|
|
|
|
imgs.append(p)
|
|
|
|
|
msks.append(m)
|
|
|
|
|
return imgs, msks
|
|
|
|
|
|
2025-08-07 18:23:53 +00:00
|
|
|
class ROISegDataset(Dataset):
|
2025-09-15 10:23:07 +00:00
|
|
|
"""
|
|
|
|
|
Compatível com o dataset original, mas agora aceita:
|
|
|
|
|
- root = '.../split/train' (com 'group' dentro)
|
|
|
|
|
- root = '.../split/train/group'
|
|
|
|
|
- root = '.../split/train/<grupo>' (ainda funciona via legado se tiver images/masks)
|
|
|
|
|
- root legado = '.../split/train' com 'images' e 'masks' diretamente
|
|
|
|
|
"""
|
2025-08-07 18:23:53 +00:00
|
|
|
def __init__(self, root, out_dir, zona_inicio, faixa_atuacao, input_w=384, min_input_h=96, labelmap_path="labelmap.txt",
|
|
|
|
|
mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)):
|
|
|
|
|
|
|
|
|
|
self.out_dir = out_dir
|
|
|
|
|
self.zona_inicio = zona_inicio
|
|
|
|
|
self.faixa_atuacao = faixa_atuacao
|
|
|
|
|
self.input_w = input_w
|
|
|
|
|
self.min_input_h = min_input_h
|
|
|
|
|
self.mean = np.array(mean, dtype=np.float32).reshape(1, 1, 3)
|
|
|
|
|
self.std = np.array(std, dtype=np.float32).reshape(1, 1, 3)
|
|
|
|
|
|
2025-09-15 10:23:07 +00:00
|
|
|
# tenta agrupar; se não encontrar, cai pro legado
|
|
|
|
|
imgs, msks = _collect_pairs_grouped(root)
|
|
|
|
|
if not imgs:
|
|
|
|
|
imgs, msks = _collect_pairs_legacy(root)
|
|
|
|
|
|
|
|
|
|
assert len(imgs) == len(msks) and len(imgs) > 0, f"Nenhuma imagem/máscara encontrada em {root}"
|
|
|
|
|
self.img_paths = imgs
|
|
|
|
|
self.msk_paths = msks
|
|
|
|
|
|
|
|
|
|
# labelmap
|
2025-08-07 18:23:53 +00:00
|
|
|
_, _, self.classes, self.ignore_rgb = carregar_labelmap_completo(labelmap_path)
|
2025-09-15 10:23:07 +00:00
|
|
|
# tipicamente ignore_rgb é (255,255,255)
|
|
|
|
|
self.ignore_id = int(self.ignore_rgb[0]) if isinstance(self.ignore_rgb, (list, tuple)) else int(self.ignore_rgb)
|
2025-08-07 18:23:53 +00:00
|
|
|
|
|
|
|
|
def __len__(self):
|
|
|
|
|
return len(self.img_paths)
|
|
|
|
|
|
|
|
|
|
def __getitem__(self, idx):
|
|
|
|
|
img_rgb = np.array(Image.open(self.img_paths[idx]).convert("RGB"))
|
2025-09-15 10:23:07 +00:00
|
|
|
# Máscara como escala de cinza (IDs já foram normalizados na etapa de normalize)
|
2025-08-07 18:23:53 +00:00
|
|
|
msk_grayscale = np.array(Image.open(self.msk_paths[idx]).convert("L"))
|
|
|
|
|
|
|
|
|
|
H, W = img_rgb.shape[:2]
|
|
|
|
|
y_fim, y_inicio = compute_roi_indices(H, self.zona_inicio, self.faixa_atuacao)
|
|
|
|
|
|
|
|
|
|
img_roi = img_rgb[y_fim:y_inicio, 0:W]
|
|
|
|
|
msk_roi = msk_grayscale[y_fim:y_inicio, 0:W]
|
|
|
|
|
|
2025-08-08 20:09:17 +00:00
|
|
|
img_in = resize_keep_width(img_roi, self.input_w, self.min_input_h, cv2.INTER_AREA)
|
|
|
|
|
msk_ids = resize_keep_width(msk_roi, self.input_w, self.min_input_h, cv2.INTER_NEAREST)
|
2025-08-07 18:23:53 +00:00
|
|
|
|
|
|
|
|
valores_validos = list(range(len(self.classes))) + [255]
|
2025-09-15 10:23:07 +00:00
|
|
|
msk_ids[np.isin(msk_ids, valores_validos, invert=True)] = self.ignore_id
|
2025-08-07 18:23:53 +00:00
|
|
|
|
2025-09-15 10:23:07 +00:00
|
|
|
if np.all(msk_ids == self.ignore_id):
|
|
|
|
|
raise ValueError(f"Máscara {self.msk_paths[idx]} está só com valor de ignore ({self.ignore_id})")
|
2025-08-07 18:23:53 +00:00
|
|
|
|
|
|
|
|
img_f = img_in.astype(np.float32) / 255.0
|
|
|
|
|
img_f = (img_f - self.mean) / self.std
|
|
|
|
|
img_chw = np.transpose(img_f, (2, 0, 1))
|
|
|
|
|
|
|
|
|
|
if idx == 0:
|
2025-09-15 10:23:07 +00:00
|
|
|
os.makedirs(self.out_dir, exist_ok=True)
|
|
|
|
|
cv2.imwrite(os.path.join(self.out_dir, "debug_roi_input.png"), cv2.cvtColor(img_roi, cv2.COLOR_RGB2BGR))
|
2025-08-07 18:23:53 +00:00
|
|
|
cv2.imwrite(os.path.join(self.out_dir, "debug_roi_mask.png"), msk_roi)
|
|
|
|
|
|
|
|
|
|
return torch.from_numpy(img_chw).float(), torch.from_numpy(msk_ids.astype(np.int64))
|