#!/usr/bin/env python3 # -*- coding: utf-8 -*- """ Split estratificado por GRUPO com val/test só do ORIGINAL e garantia de NÃO VAZAMENTO. Lê de: MODELO/dataset//group//{images,masks,(masks2),(labels)} Escreve em: MODELO/dataset/split//group//{images,masks,(masks2),(labels)} Definições: - Família = todas as variações da mesma base original: original_.* augmented__aug_XX.* - Val/Test: somente original_. - Train: original_ + todos augmented__aug_XX. Labels: - Ativados por config['dual_head_label']. - Copia labels .json/.txt e .npy quando existirem. - O pareamento é feito pelo mesmo base name da imagem. Uso: python _7_split_grouped_noleak_with_labels.py python _7_split_grouped_noleak_with_labels.py --train 0.7 --val 0.29 --test 0.01 --seed 42 python _7_split_grouped_noleak_with_labels.py --strict-label python _7_split_grouped_noleak_with_labels.py --cap-train-families "navegavel:300,naonavegavel_navegavel:800" """ import os import re import json import shutil import random import argparse from typing import Dict, List, Optional # ===================== CONFIG ===================== with open("config.json", "r", encoding="utf-8") as f: config = json.load(f) MODELO = config.get("camera") USE_MASKS2 = bool(config.get("dual_head_mask", config.get("dual_head", False))) USE_LABELS = bool(config.get("dual_head_label", False)) RESOLUCAO = tuple(config.get("resolucao")) pasta_origem = os.path.join(MODELO, "dataset", f"{RESOLUCAO[0]}x{RESOLUCAO[1]}", "group") pasta_destino = os.path.join(MODELO, "dataset", "split") IMG_EXTS = (".jpg", ".jpeg", ".png") MSK_EXT = ".png" LABEL_EXTS = (".json", ".txt", ".npy") RE_ORIGINAL_PREFIX = re.compile(r"^original_(.+)$", re.IGNORECASE) RE_AUGMENTED_FAMILY = re.compile(r"^augmented_(.+?)(?:_aug_\d+)?$", re.IGNORECASE) RE_AUG_SUFFIX = re.compile(r"_aug_\d+$", re.IGNORECASE) # ===================== HELPERS ===================== def garantir(p): os.makedirs(p, exist_ok=True) def lista_grupos(root): if not os.path.isdir(root): return [] out = [] for g in sorted(os.listdir(root)): gdir = os.path.join(root, g) 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")): out.append(g) return out def listar_imagens(img_dir): if not os.path.isdir(img_dir): return [] fs = [] for f in os.listdir(img_dir): ext = os.path.splitext(f.lower())[1] if ext in IMG_EXTS: fs.append(f) return sorted(fs) def mask_from_image_name(img_name): base, _ = os.path.splitext(img_name) return base + MSK_EXT def mask2_from_image_name(img_name): base, _ = os.path.splitext(img_name) return base + MSK_EXT def label_candidates_from_image_name(img_name): base, _ = os.path.splitext(img_name) return [base + ext for ext in LABEL_EXTS] def classify_source_and_family(filename_no_ext): """ Retorna (source, family_key) source ∈ {original, augmented, unknown} family_key = base original sem prefixo/sufixo. """ m = RE_ORIGINAL_PREFIX.match(filename_no_ext) if m: return "original", m.group(1) m = RE_AUGMENTED_FAMILY.match(filename_no_ext) if m: return "augmented", m.group(1) if RE_AUG_SUFFIX.search(filename_no_ext): fam = RE_AUG_SUFFIX.sub("", filename_no_ext) return "augmented", fam return "unknown", filename_no_ext def has_any_label(label_dir: str, img_name: str) -> bool: if not label_dir or not os.path.isdir(label_dir): return False return any(os.path.exists(os.path.join(label_dir, cand)) for cand in label_candidates_from_image_name(img_name)) def build_family_index(img_dir, msk_dir, label_dir=None, require_label=False): """ Constrói índice de famílias a partir de img_dir/msk_dir/labels. Retorna: dict family -> { original: str|None, augmented: [str], all: [str] } Os nomes são arquivos de imagem. """ familias = {} imgs = listar_imagens(img_dir) for img_name in imgs: base_no_ext, _ = os.path.splitext(img_name) mask_name = mask_from_image_name(img_name) if not os.path.exists(os.path.join(msk_dir, mask_name)): continue if require_label and not has_any_label(label_dir, img_name): continue source, fam = classify_source_and_family(base_no_ext) d = familias.setdefault(fam, {"original": None, "augmented": [], "all": []}) d["all"].append(img_name) if source == "original": d["original"] = img_name elif source == "augmented": d["augmented"].append(img_name) else: if d["original"] is None: d["original"] = img_name else: d["augmented"].append(img_name) return familias def allocate_counts(n, p_train, p_val, p_test, min_train, min_val, min_test): n_train = int(round(n * p_train)) n_val = int(round(n * p_val)) n_test = n - n_train - n_val if n_test < 0: excesso = -n_test take_train = min(excesso, max(0, n_train)) n_train -= take_train excesso -= take_train if excesso > 0: take_val = min(excesso, max(0, n_val)) n_val -= take_val excesso -= take_val n_test = 0 min_sum = min_train + min_val + min_test if n >= min_sum: n_train = max(n_train, min_train) n_val = max(n_val, min_val) n_test = max(n_test, min_test) total = n_train + n_val + n_test while total > n: if n_test > min_test: n_test -= 1 elif n_val > min_val: n_val -= 1 elif n_train > min_train: n_train -= 1 else: break total = n_train + n_val + n_test while total < n: if n_train - min_train <= n_val - min_val: n_train += 1 else: n_val += 1 total = n_train + n_val + n_test else: n_train = min(n, max(1, min_train)) resto = n - n_train n_val = max(0, min(resto, min_val)) n_test = max(0, resto - n_val) diff = n - (n_train + n_val + n_test) if diff != 0: if diff > 0: take = min(diff, n - n_train) n_train += take diff -= take if diff > 0: n_val += diff else: diff = -diff take = min(diff, n_test) n_test -= take diff -= take if diff > 0: n_val -= diff return n_train, n_val, n_test def copiar_labels_para_item(nome, src_label_dir, dst_label_dir): if not src_label_dir or not dst_label_dir or not os.path.isdir(src_label_dir): return 0 garantir(dst_label_dir) copied = 0 for cand in label_candidates_from_image_name(nome): src = os.path.join(src_label_dir, cand) if not os.path.exists(src): continue shutil.copy2(src, os.path.join(dst_label_dir, cand)) copied += 1 return copied def copiar( nomes, src_img_dir, src_msk_dir, dst_img_dir, dst_msk_dir, src_msk2_dir=None, dst_msk2_dir=None, src_label_dir=None, dst_label_dir=None, strict_label=False, ): garantir(dst_img_dir) garantir(dst_msk_dir) use_msk2 = bool(src_msk2_dir and dst_msk2_dir and os.path.isdir(src_msk2_dir)) use_labels = bool(src_label_dir and dst_label_dir and os.path.isdir(src_label_dir)) if use_msk2: garantir(dst_msk2_dir) if use_labels: garantir(dst_label_dir) moved = 0 skipped_no_label = 0 for nome in nomes: mask_name = mask_from_image_name(nome) src_img = os.path.join(src_img_dir, nome) src_msk = os.path.join(src_msk_dir, mask_name) if not (os.path.exists(src_img) and os.path.exists(src_msk)): continue if strict_label and use_labels and not has_any_label(src_label_dir, nome): skipped_no_label += 1 continue shutil.copy2(src_img, os.path.join(dst_img_dir, nome)) shutil.copy2(src_msk, os.path.join(dst_msk_dir, mask_name)) if use_msk2: m2_name = mask2_from_image_name(nome) src_m2 = os.path.join(src_msk2_dir, m2_name) if os.path.exists(src_m2): shutil.copy2(src_m2, os.path.join(dst_msk2_dir, m2_name)) if use_labels: copied = copiar_labels_para_item(nome, src_label_dir, dst_label_dir) if strict_label and copied == 0: skipped_no_label += 1 continue moved += 1 if skipped_no_label > 0: print(f"[WARN] {skipped_no_label} itens pulados por falta de label.") return moved # ===================== SPLIT ===================== def split_group(group_name, p_train, p_val, p_test, seed, mins, caps_map=None, strict_label=False): src_img_dir = os.path.join(pasta_origem, group_name, "images") src_msk_dir = os.path.join(pasta_origem, group_name, "masks") src_msk2_dir = os.path.join(pasta_origem, group_name, "masks2") src_label_dir = os.path.join(pasta_origem, group_name, "labels") use_msk2 = USE_MASKS2 and os.path.isdir(src_msk2_dir) use_labels = USE_LABELS and os.path.isdir(src_label_dir) if USE_LABELS and not use_labels: msg = f"[{group_name}] dual_head_label=true, mas labels/ não existe." if strict_label: print(f"[WARN] {msg} Pulando grupo.") return {"train": 0, "val": 0, "test": 0, "familias": 0} print(f"[WARN] {msg} Split seguirá sem copiar labels.") familias = build_family_index( src_img_dir, src_msk_dir, label_dir=src_label_dir if use_labels else None, require_label=bool(strict_label and use_labels), ) familias_originais = [fam for fam, d in familias.items() if d["original"] is not None] total_familias = len(familias_originais) if total_familias == 0: print(f"[{group_name}] 0 famílias com original, pulando.") return {"train": 0, "val": 0, "test": 0, "familias": 0} rng = random.Random(seed) rng.shuffle(familias_originais) n_tr, n_va, n_te = allocate_counts( total_familias, p_train, p_val, p_test, mins["train"], mins["val"], mins["test"], ) fam_train = set(familias_originais[:n_tr]) fam_val = set(familias_originais[n_tr : n_tr + n_va]) fam_test = set(familias_originais[n_tr + n_va : n_tr + n_va + n_te]) if caps_map and group_name in caps_map: cap = caps_map[group_name] if len(fam_train) > cap: fam_list = list(fam_train) rng.shuffle(fam_list) kept = set(fam_list[:cap]) dropped = set(fam_list[cap:]) fam_train = kept print( f"[{group_name}] cap-train-families={cap} → " f"mantidas {len(kept)} famílias, descartadas {len(dropped)} do TRAIN" ) nomes_train, nomes_val, nomes_test = [], [], [] for fam, d in familias.items(): if fam in fam_train: if d["original"]: nomes_train.append(d["original"]) if d["augmented"]: nomes_train.extend(d["augmented"]) elif fam in fam_val: if d["original"]: nomes_val.append(d["original"]) elif fam in fam_test: if d["original"]: nomes_test.append(d["original"]) dest_train_img = os.path.join(pasta_destino, "train", "group", group_name, "images") dest_train_msk = os.path.join(pasta_destino, "train", "group", group_name, "masks") dest_val_img = os.path.join(pasta_destino, "val", "group", group_name, "images") dest_val_msk = os.path.join(pasta_destino, "val", "group", group_name, "masks") dest_test_img = os.path.join(pasta_destino, "test", "group", group_name, "images") dest_test_msk = os.path.join(pasta_destino, "test", "group", group_name, "masks") dest_train_msk2 = os.path.join(pasta_destino, "train", "group", group_name, "masks2") if use_msk2 else None dest_val_msk2 = os.path.join(pasta_destino, "val", "group", group_name, "masks2") if use_msk2 else None dest_test_msk2 = os.path.join(pasta_destino, "test", "group", group_name, "masks2") if use_msk2 else None dest_train_label = os.path.join(pasta_destino, "train", "group", group_name, "labels") if use_labels else None dest_val_label = os.path.join(pasta_destino, "val", "group", group_name, "labels") if use_labels else None dest_test_label = os.path.join(pasta_destino, "test", "group", group_name, "labels") if use_labels else None m_train = copiar( nomes_train, src_img_dir, src_msk_dir, dest_train_img, dest_train_msk, src_msk2_dir, dest_train_msk2, src_label_dir, dest_train_label, strict_label=strict_label, ) m_val = copiar( nomes_val, src_img_dir, src_msk_dir, dest_val_img, dest_val_msk, src_msk2_dir, dest_val_msk2, src_label_dir, dest_val_label, strict_label=strict_label, ) m_test = copiar( nomes_test, src_img_dir, src_msk_dir, dest_test_img, dest_test_msk, src_msk2_dir, dest_test_msk2, src_label_dir, dest_test_label, strict_label=strict_label, ) print( f"[{group_name}] famílias={total_familias} → " f"train(imgs)={m_train}, val(imgs)={m_val}, test(imgs)={m_test}" ) return {"train": m_train, "val": m_val, "test": m_test, "familias": total_familias} # ===================== MAIN ===================== def parse_cap_map(s): caps = {} if not s: return caps for item in s.split(","): if not item.strip(): continue k, v = item.strip().split(":") caps[k.strip()] = int(v) return caps def main(): ap = argparse.ArgumentParser(description="Split estratificado por grupo sem vazamento, com suporte a labels.") ap.add_argument("--train", type=float, default=0.70, help="Proporção de treino.") ap.add_argument("--val", type=float, default=0.29, help="Proporção de validação.") ap.add_argument("--test", type=float, default=0.01, help="Proporção de teste.") ap.add_argument("--seed", type=int, default=42, help="Seed do embaralhamento.") ap.add_argument("--min-train", type=int, default=1, help="Mínimo de famílias por grupo em train.") ap.add_argument("--min-val", type=int, default=1, help="Mínimo de famílias por grupo em val.") ap.add_argument("--min-test", type=int, default=0, help="Mínimo de famílias por grupo em test.") ap.add_argument("--modelo", type=str, default=None, help="Sobrescreve MODELO do config.json.") ap.add_argument("--resolucao", type=str, default=None, help="Sobrescreve resolução no formato WxH.") ap.add_argument( "--cap-train-families", type=str, default="", help="Mapa grupo:cap para limitar famílias no TRAIN. Ex: 'navegavel:350'", ) ap.add_argument( "--strict-label", action="store_true", help="Se dual_head_label=true e faltar label, pula item/grupo.", ) args = ap.parse_args() modelo = args.modelo or MODELO if args.resolucao: try: w, h = args.resolucao.lower().split("x") resolucao = (int(w), int(h)) except Exception: resolucao = RESOLUCAO else: resolucao = RESOLUCAO caps_map = parse_cap_map(args.cap_train_families) global pasta_origem, pasta_destino pasta_origem = os.path.join(modelo, "dataset", f"{resolucao[0]}x{resolucao[1]}", "group") pasta_destino = os.path.join(modelo, "dataset", "split") soma = args.train + args.val + args.test if soma <= 0: raise ValueError("Soma de proporções deve ser > 0.") p_train = args.train / soma p_val = args.val / soma p_test = args.test / soma mins = { "train": max(0, args.min_train), "val": max(0, args.min_val), "test": max(0, args.min_test), } garantir(pasta_destino) grupos = lista_grupos(pasta_origem) if not grupos: print(f"[WARN] Nenhum grupo encontrado em: {pasta_origem}") return random.seed(args.seed) total_global = {"train": 0, "val": 0, "test": 0, "familias": 0} print(f"[INFO] Origem: {pasta_origem}") print(f"[INFO] Destino: {pasta_destino}") print(f"[INFO] dual_head_mask/masks2: {USE_MASKS2}") print(f"[INFO] dual_head_label/labels: {USE_LABELS}") print(f"Grupos: {', '.join(grupos)}") print(f"Proporções normalizadas: train={p_train:.3f}, val={p_val:.3f}, test={p_test:.3f}") print(f"Mínimos por grupo: train={mins['train']} val={mins['val']} test={mins['test']}") for g in grupos: res = split_group(g, p_train, p_val, p_test, args.seed, mins, caps_map=caps_map, strict_label=args.strict_label) for k in total_global.keys(): total_global[k] += res.get(k, 0) print("\nResumo global, imagens copiadas:") print(f" train: {total_global['train']}") print(f" val: {total_global['val']}") print(f" test: {total_global['test']}") print(f" famílias total: {total_global['familias']}") print("\n✅ Split sem vazamento concluído!") if __name__ == "__main__": main()