# -*- coding: utf-8 -*- 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 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//{images,masks} - root//{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 class ROISegDataset(Dataset): """ Compatível com o dataset original, mas agora aceita: - root = '.../split/train' (com 'group' dentro) - root = '.../split/train/group' - root = '.../split/train/' (ainda funciona via legado se tiver images/masks) - root legado = '.../split/train' com 'images' e 'masks' diretamente """ 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) # 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 _, _, self.classes, self.ignore_rgb = carregar_labelmap_completo(labelmap_path) # 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) def __len__(self): return len(self.img_paths) def __getitem__(self, idx): img_rgb = np.array(Image.open(self.img_paths[idx]).convert("RGB")) # Máscara como escala de cinza (IDs já foram normalizados na etapa de normalize) 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] 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) valores_validos = list(range(len(self.classes))) + [255] msk_ids[np.isin(msk_ids, valores_validos, invert=True)] = self.ignore_id 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})") 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: 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)) 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))