agrobot_base/Python/OAK/datasets/_7_split.py

564 lines
17 KiB
Python
Raw Normal View History

2025-09-15 10:23:07 +00:00
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
2026-05-21 22:57:35 +00:00
Split estratificado por GRUPO com val/test do ORIGINAL e garantia de NÃO VAZAMENTO.
2025-09-15 10:23:07 +00:00
de:
2026-05-21 22:57:35 +00:00
MODELO/dataset/<WxH>/group/<grupo>/{images,masks,(masks2),(labels)}
2025-09-15 10:23:07 +00:00
Escreve em:
2026-05-21 22:57:35 +00:00
MODELO/dataset/split/<split>/group/<grupo>/{images,masks,(masks2),(labels)}
2025-09-15 10:23:07 +00:00
Definições:
2026-05-21 22:57:35 +00:00
- Família = todas as variações da mesma base original:
original_<base>.*
augmented_<base>_aug_XX.*
- Val/Test: somente original_<base>.
- Train: original_<base> + todos augmented_<base>_aug_XX.
2025-09-15 10:23:07 +00:00
2026-05-21 22:57:35 +00:00
Labels:
- Ativados por config['dual_head_label'].
- Copia labels .json/.txt e .npy quando existirem.
- O pareamento é feito pelo mesmo base name da imagem.
2025-09-15 10:23:07 +00:00
Uso:
2026-05-21 22:57:35 +00:00
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"
2025-09-15 10:23:07 +00:00
"""
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
import os
import re
import json
import shutil
import random
import argparse
2026-05-21 22:57:35 +00:00
from typing import Dict, List, Optional
# ===================== CONFIG =====================
2025-09-15 10:23:07 +00:00
2026-05-21 22:57:35 +00:00
with open("config.json", "r", encoding="utf-8") as f:
2025-09-15 10:23:07 +00:00
config = json.load(f)
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
MODELO = config.get("camera")
2026-05-21 22:57:35 +00:00
USE_MASKS2 = bool(config.get("dual_head_mask", config.get("dual_head", False)))
USE_LABELS = bool(config.get("dual_head_label", False))
2025-09-15 10:23:07 +00:00
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")
2026-05-21 22:57:35 +00:00
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)
2025-09-15 10:23:07 +00:00
2026-05-21 22:57:35 +00:00
# ===================== HELPERS =====================
2025-09-15 10:23:07 +00:00
def garantir(p):
os.makedirs(p, exist_ok=True)
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
def lista_grupos(root):
if not os.path.isdir(root):
return []
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
out = []
for g in sorted(os.listdir(root)):
gdir = os.path.join(root, g)
2026-05-21 22:57:35 +00:00
if not os.path.isdir(gdir):
continue
2025-09-15 10:23:07 +00:00
if os.path.isdir(os.path.join(gdir, "images")) and os.path.isdir(os.path.join(gdir, "masks")):
out.append(g)
return out
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
def listar_imagens(img_dir):
2026-05-21 22:57:35 +00:00
if not os.path.isdir(img_dir):
return []
2025-09-15 10:23:07 +00:00
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)
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
def mask_from_image_name(img_name):
base, _ = os.path.splitext(img_name)
return base + MSK_EXT
2026-05-21 22:57:35 +00:00
def mask2_from_image_name(img_name):
base, _ = os.path.splitext(img_name)
2026-05-21 22:57:35 +00:00
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]
2025-09-15 10:23:07 +00:00
def classify_source_and_family(filename_no_ext):
"""
Retorna (source, family_key)
2026-05-21 22:57:35 +00:00
source {original, augmented, unknown}
family_key = base original sem prefixo/sufixo.
2025-09-15 10:23:07 +00:00
"""
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
2026-05-21 22:57:35 +00:00
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):
2025-09-15 10:23:07 +00:00
"""
2026-05-21 22:57:35 +00:00
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.
2025-09-15 10:23:07 +00:00
"""
familias = {}
imgs = listar_imagens(img_dir)
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
for img_name in imgs:
2026-05-21 22:57:35 +00:00
base_no_ext, _ = os.path.splitext(img_name)
2025-09-15 10:23:07 +00:00
mask_name = mask_from_image_name(img_name)
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
if not os.path.exists(os.path.join(msk_dir, mask_name)):
2026-05-21 22:57:35 +00:00
continue
if require_label and not has_any_label(label_dir, img_name):
continue
2025-09-15 10:23:07 +00:00
source, fam = classify_source_and_family(base_no_ext)
d = familias.setdefault(fam, {"original": None, "augmented": [], "all": []})
d["all"].append(img_name)
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
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)
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
return familias
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
def allocate_counts(n, p_train, p_val, p_test, min_train, min_val, min_test):
n_train = int(round(n * p_train))
2026-05-21 22:57:35 +00:00
n_val = int(round(n * p_val))
n_test = n - n_train - n_val
2025-09-15 10:23:07 +00:00
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)
2026-05-21 22:57:35 +00:00
n_val = max(n_val, min_val)
n_test = max(n_test, min_test)
2025-09-15 10:23:07 +00:00
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
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
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
2026-05-21 22:57:35 +00:00
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))
2026-05-21 22:57:35 +00:00
use_labels = bool(src_label_dir and dst_label_dir and os.path.isdir(src_label_dir))
if use_msk2:
garantir(dst_msk2_dir)
2026-05-21 22:57:35 +00:00
if use_labels:
garantir(dst_label_dir)
2025-09-15 10:23:07 +00:00
moved = 0
2026-05-21 22:57:35 +00:00
skipped_no_label = 0
2025-09-15 10:23:07 +00:00
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)
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
if not (os.path.exists(src_img) and os.path.exists(src_msk)):
continue
2026-05-21 22:57:35 +00:00
if strict_label and use_labels and not has_any_label(src_label_dir, nome):
skipped_no_label += 1
continue
2025-09-15 10:23:07 +00:00
shutil.copy2(src_img, os.path.join(dst_img_dir, nome))
shutil.copy2(src_msk, os.path.join(dst_msk_dir, mask_name))
2026-05-21 22:57:35 +00:00
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))
2026-05-21 22:57:35 +00:00
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
2025-09-15 10:23:07 +00:00
moved += 1
2026-05-21 22:57:35 +00:00
if skipped_no_label > 0:
print(f"[WARN] {skipped_no_label} itens pulados por falta de label.")
2025-09-15 10:23:07 +00:00
return moved
2026-05-21 22:57:35 +00:00
# ===================== SPLIT =====================
def split_group(group_name, p_train, p_val, p_test, seed, mins, caps_map=None, strict_label=False):
2025-09-15 10:23:07 +00:00
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")
2026-05-21 22:57:35 +00:00
src_label_dir = os.path.join(pasta_origem, group_name, "labels")
use_msk2 = USE_MASKS2 and os.path.isdir(src_msk2_dir)
2026-05-21 22:57:35 +00:00
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),
)
2025-09-15 10:23:07 +00:00
familias_originais = [fam for fam, d in familias.items() if d["original"] is not None]
total_familias = len(familias_originais)
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
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(
2026-05-21 22:57:35 +00:00
total_familias,
p_train,
p_val,
p_test,
mins["train"],
mins["val"],
mins["test"],
2025-09-15 10:23:07 +00:00
)
fam_train = set(familias_originais[:n_tr])
2026-05-21 22:57:35 +00:00
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])
2025-09-15 10:23:07 +00:00
if caps_map and group_name in caps_map:
cap = caps_map[group_name]
if len(fam_train) > cap:
fam_list = list(fam_train)
2026-05-21 22:57:35 +00:00
rng.shuffle(fam_list)
kept = set(fam_list[:cap])
dropped = set(fam_list[cap:])
2025-09-15 10:23:07 +00:00
fam_train = kept
2026-05-21 22:57:35 +00:00
print(
f"[{group_name}] cap-train-families={cap}"
f"mantidas {len(kept)} famílias, descartadas {len(dropped)} do TRAIN"
)
2025-09-15 10:23:07 +00:00
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")
2026-05-21 22:57:35 +00:00
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")
2025-09-15 10:23:07 +00:00
dest_train_msk2 = os.path.join(pasta_destino, "train", "group", group_name, "masks2") if use_msk2 else None
2026-05-21 22:57:35 +00:00
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,
)
2025-09-15 10:23:07 +00:00
2026-05-21 22:57:35 +00:00
print(
f"[{group_name}] famílias={total_familias}"
f"train(imgs)={m_train}, val(imgs)={m_val}, test(imgs)={m_test}"
)
2025-09-15 10:23:07 +00:00
return {"train": m_train, "val": m_val, "test": m_test, "familias": total_familias}
2026-05-21 22:57:35 +00:00
# ===================== 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
2025-09-15 10:23:07 +00:00
def main():
2026-05-21 22:57:35 +00:00
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.")
2025-09-15 10:23:07 +00:00
2026-05-21 22:57:35 +00:00
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.")
2025-09-15 10:23:07 +00:00
ap.add_argument("--modelo", type=str, default=None, help="Sobrescreve MODELO do config.json.")
2026-05-21 22:57:35 +00:00
ap.add_argument("--resolucao", type=str, default=None, help="Sobrescreve resolução no formato WxH.")
2025-09-15 10:23:07 +00:00
2026-05-21 22:57:35 +00:00
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.",
)
2025-09-15 10:23:07 +00:00
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
2026-05-21 22:57:35 +00:00
if soma <= 0:
raise ValueError("Soma de proporções deve ser > 0.")
2025-09-15 10:23:07 +00:00
p_train = args.train / soma
2026-05-21 22:57:35 +00:00
p_val = args.val / soma
p_test = args.test / soma
2025-09-15 10:23:07 +00:00
2026-05-21 22:57:35 +00:00
mins = {
"train": max(0, args.min_train),
"val": max(0, args.min_val),
"test": max(0, args.min_test),
}
2025-09-15 10:23:07 +00:00
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)
2026-05-21 22:57:35 +00:00
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}")
2025-09-15 10:23:07 +00:00
print(f"Grupos: {', '.join(grupos)}")
print(f"Proporções normalizadas: train={p_train:.3f}, val={p_val:.3f}, test={p_test:.3f}")
2026-05-21 22:57:35 +00:00
print(f"Mínimos por grupo: train={mins['train']} val={mins['val']} test={mins['test']}")
2025-09-15 10:23:07 +00:00
for g in grupos:
2026-05-21 22:57:35 +00:00
res = split_group(g, p_train, p_val, p_test, args.seed, mins, caps_map=caps_map, strict_label=args.strict_label)
2025-09-15 10:23:07 +00:00
for k in total_global.keys():
total_global[k] += res.get(k, 0)
2026-05-21 22:57:35 +00:00
print("\nResumo global, imagens copiadas:")
2025-09-15 10:23:07 +00:00
print(f" train: {total_global['train']}")
print(f" val: {total_global['val']}")
print(f" test: {total_global['test']}")
2026-05-21 22:57:35 +00:00
print(f" famílias total: {total_global['familias']}")
2025-09-15 10:23:07 +00:00
print("\n✅ Split sem vazamento concluído!")
2026-05-21 22:57:35 +00:00
2025-09-15 10:23:07 +00:00
if __name__ == "__main__":
main()