2026-01-22 18:47:56 +00:00
|
|
|
#!/usr/bin/env python3
|
|
|
|
|
# -*- coding: utf-8 -*-
|
|
|
|
|
|
|
|
|
|
import os
|
|
|
|
|
import re
|
|
|
|
|
import json
|
|
|
|
|
import shutil
|
|
|
|
|
import random
|
|
|
|
|
import argparse
|
|
|
|
|
|
2026-04-22 20:08:49 +00:00
|
|
|
|
2026-01-22 18:47:56 +00:00
|
|
|
with open("config.json", "r", encoding="utf-8") as f:
|
|
|
|
|
config = json.load(f)
|
|
|
|
|
|
2026-04-22 20:08:49 +00:00
|
|
|
RESOLUCAO = tuple(config.get("resolucao"))
|
2026-04-20 18:53:28 +00:00
|
|
|
pasta_origem = os.path.join("dataset", f"{RESOLUCAO[0]}x{RESOLUCAO[1]}", "group")
|
|
|
|
|
pasta_destino = os.path.join("dataset", "split")
|
2026-01-22 18:47:56 +00:00
|
|
|
|
2026-04-22 20:08:49 +00:00
|
|
|
TENSOR_EXT = ".npy"
|
|
|
|
|
MASK_NPY_SUFFIX = ".npy"
|
2026-01-22 18:47:56 +00:00
|
|
|
|
2026-04-22 20:08:49 +00:00
|
|
|
RE_ORIGINAL_PREFIX = re.compile(r"^original_(.+)$", re.IGNORECASE)
|
|
|
|
|
RE_AUGMENTED_FAMILY = re.compile(r"^augmented_(.+?)(?:_aug[a-zA-Z0-9]*_\d+)?$", re.IGNORECASE)
|
|
|
|
|
RE_AUG_SUFFIX = re.compile(r"_aug[a-zA-Z0-9]*_\d+$", re.IGNORECASE)
|
2026-01-22 18:47:56 +00:00
|
|
|
|
|
|
|
|
|
|
|
|
|
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
|
2026-04-22 20:08:49 +00:00
|
|
|
if os.path.isdir(os.path.join(gdir, "tensors")) and os.path.isdir(os.path.join(gdir, "masks")):
|
2026-01-22 18:47:56 +00:00
|
|
|
out.append(g)
|
|
|
|
|
return out
|
|
|
|
|
|
|
|
|
|
|
2026-04-22 20:08:49 +00:00
|
|
|
def listar_tensors(tensor_dir):
|
|
|
|
|
if not os.path.isdir(tensor_dir):
|
2026-01-22 18:47:56 +00:00
|
|
|
return []
|
|
|
|
|
fs = []
|
2026-04-22 20:08:49 +00:00
|
|
|
for f in os.listdir(tensor_dir):
|
|
|
|
|
if f.lower().endswith(TENSOR_EXT):
|
2026-01-22 18:47:56 +00:00
|
|
|
fs.append(f)
|
|
|
|
|
return sorted(fs)
|
|
|
|
|
|
|
|
|
|
|
2026-04-22 20:08:49 +00:00
|
|
|
def mask_npy_from_tensor_name(tensor_name):
|
|
|
|
|
base, _ = os.path.splitext(tensor_name)
|
|
|
|
|
return base + MASK_NPY_SUFFIX
|
2026-01-22 18:47:56 +00:00
|
|
|
|
|
|
|
|
|
|
|
|
|
def classify_source_and_family(filename_no_ext):
|
|
|
|
|
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-04-22 20:08:49 +00:00
|
|
|
def build_family_index(tensor_dir, mask_dir):
|
2026-01-22 18:47:56 +00:00
|
|
|
"""
|
2026-04-22 20:08:49 +00:00
|
|
|
family -> {
|
|
|
|
|
"original": tensor_name or None,
|
|
|
|
|
"augmented": [tensor_name, ...],
|
|
|
|
|
"all": [...]
|
|
|
|
|
}
|
|
|
|
|
Só indexa se houver pelo menos mask .npy correspondente.
|
2026-01-22 18:47:56 +00:00
|
|
|
"""
|
|
|
|
|
familias = {}
|
2026-04-22 20:08:49 +00:00
|
|
|
tensors = listar_tensors(tensor_dir)
|
|
|
|
|
|
|
|
|
|
for tensor_name in tensors:
|
|
|
|
|
base_no_ext, _ = os.path.splitext(tensor_name)
|
|
|
|
|
mask_npy_name = mask_npy_from_tensor_name(tensor_name)
|
|
|
|
|
|
|
|
|
|
if not os.path.exists(os.path.join(mask_dir, mask_npy_name)):
|
|
|
|
|
continue
|
2026-01-22 18:47:56 +00:00
|
|
|
|
|
|
|
|
source, fam = classify_source_and_family(base_no_ext)
|
|
|
|
|
d = familias.setdefault(fam, {"original": None, "augmented": [], "all": []})
|
2026-04-22 20:08:49 +00:00
|
|
|
d["all"].append(tensor_name)
|
|
|
|
|
|
2026-01-22 18:47:56 +00:00
|
|
|
if source == "original":
|
2026-04-22 20:08:49 +00:00
|
|
|
d["original"] = tensor_name
|
2026-01-22 18:47:56 +00:00
|
|
|
elif source == "augmented":
|
2026-04-22 20:08:49 +00:00
|
|
|
d["augmented"].append(tensor_name)
|
2026-01-22 18:47:56 +00:00
|
|
|
else:
|
|
|
|
|
if d["original"] is None:
|
2026-04-22 20:08:49 +00:00
|
|
|
d["original"] = tensor_name
|
2026-01-22 18:47:56 +00:00
|
|
|
else:
|
2026-04-22 20:08:49 +00:00
|
|
|
d["augmented"].append(tensor_name)
|
|
|
|
|
|
2026-01-22 18:47:56 +00:00
|
|
|
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))
|
2026-04-22 20:08:49 +00:00
|
|
|
n_val = int(round(n * p_val))
|
|
|
|
|
n_test = n - n_train - n_val
|
2026-01-22 18:47:56 +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-04-22 20:08:49 +00:00
|
|
|
n_val = max(n_val, min_val)
|
|
|
|
|
n_test = max(n_test, min_test)
|
2026-01-22 18:47:56 +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-04-22 20:08:49 +00:00
|
|
|
|
2026-01-22 18:47:56 +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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def copiar(
|
|
|
|
|
nomes,
|
2026-04-22 20:08:49 +00:00
|
|
|
src_tensor_dir,
|
|
|
|
|
src_mask_dir,
|
|
|
|
|
dst_tensor_dir,
|
|
|
|
|
dst_mask_dir,
|
2026-01-22 18:47:56 +00:00
|
|
|
):
|
|
|
|
|
"""
|
2026-04-22 20:08:49 +00:00
|
|
|
Copia:
|
|
|
|
|
- tensor .npy
|
|
|
|
|
- mask .npy obrigatória
|
|
|
|
|
- mask .png opcional (debug)
|
2026-01-22 18:47:56 +00:00
|
|
|
"""
|
2026-04-22 20:08:49 +00:00
|
|
|
garantir(dst_tensor_dir)
|
|
|
|
|
garantir(dst_mask_dir)
|
2026-01-22 18:47:56 +00:00
|
|
|
|
|
|
|
|
moved = 0
|
|
|
|
|
for nome in nomes:
|
2026-04-22 20:08:49 +00:00
|
|
|
tensor_src = os.path.join(src_tensor_dir, nome)
|
|
|
|
|
mask_npy_name = mask_npy_from_tensor_name(nome)
|
2026-01-22 18:47:56 +00:00
|
|
|
|
2026-04-22 20:08:49 +00:00
|
|
|
mask_npy_src = os.path.join(src_mask_dir, mask_npy_name)
|
2026-01-22 18:47:56 +00:00
|
|
|
|
2026-04-22 20:08:49 +00:00
|
|
|
if not (os.path.exists(tensor_src) and os.path.exists(mask_npy_src)):
|
|
|
|
|
continue
|
2026-01-22 18:47:56 +00:00
|
|
|
|
2026-04-22 20:08:49 +00:00
|
|
|
shutil.copy2(tensor_src, os.path.join(dst_tensor_dir, nome))
|
|
|
|
|
shutil.copy2(mask_npy_src, os.path.join(dst_mask_dir, mask_npy_name))
|
2026-01-22 18:47:56 +00:00
|
|
|
|
|
|
|
|
moved += 1
|
2026-04-22 20:08:49 +00:00
|
|
|
|
2026-01-22 18:47:56 +00:00
|
|
|
return moved
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def split_group(group_name, p_train, p_val, p_test, seed, mins, caps_map=None):
|
2026-04-22 20:08:49 +00:00
|
|
|
src_tensor_dir = os.path.join(pasta_origem, group_name, "tensors")
|
|
|
|
|
src_mask_dir = os.path.join(pasta_origem, group_name, "masks")
|
2026-01-22 18:47:56 +00:00
|
|
|
|
2026-04-22 20:08:49 +00:00
|
|
|
familias = build_family_index(src_tensor_dir, src_mask_dir)
|
2026-01-22 18:47:56 +00:00
|
|
|
|
|
|
|
|
familias_originais = [fam for fam, d in familias.items() if d["original"] is not None]
|
|
|
|
|
total_familias = len(familias_originais)
|
2026-04-22 20:08:49 +00:00
|
|
|
|
2026-01-22 18:47:56 +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(
|
|
|
|
|
total_familias, p_train, p_val, p_test,
|
|
|
|
|
mins["train"], mins["val"], mins["test"]
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
fam_train = set(familias_originais[:n_tr])
|
2026-04-22 20:08:49 +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])
|
2026-01-22 18:47:56 +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)
|
|
|
|
|
rng.shuffle(fam_list)
|
2026-04-22 20:08:49 +00:00
|
|
|
kept = set(fam_list[:cap])
|
2026-01-22 18:47:56 +00:00
|
|
|
dropped = set(fam_list[cap:])
|
|
|
|
|
fam_train = kept
|
|
|
|
|
print(f"[{group_name}] cap-train-families={cap} → 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"])
|
|
|
|
|
|
2026-04-22 20:08:49 +00:00
|
|
|
dst_train_tensor = os.path.join(pasta_destino, "train", "group", group_name, "tensors")
|
|
|
|
|
dst_train_mask = os.path.join(pasta_destino, "train", "group", group_name, "masks")
|
|
|
|
|
|
|
|
|
|
dst_val_tensor = os.path.join(pasta_destino, "val", "group", group_name, "tensors")
|
|
|
|
|
dst_val_mask = os.path.join(pasta_destino, "val", "group", group_name, "masks")
|
|
|
|
|
|
|
|
|
|
dst_test_tensor = os.path.join(pasta_destino, "test", "group", group_name, "tensors")
|
|
|
|
|
dst_test_mask = os.path.join(pasta_destino, "test", "group", group_name, "masks")
|
|
|
|
|
|
|
|
|
|
m_train = copiar(nomes_train, src_tensor_dir, src_mask_dir, dst_train_tensor, dst_train_mask)
|
|
|
|
|
m_val = copiar(nomes_val, src_tensor_dir, src_mask_dir, dst_val_tensor, dst_val_mask)
|
|
|
|
|
m_test = copiar(nomes_test, src_tensor_dir, src_mask_dir, dst_test_tensor, dst_test_mask)
|
|
|
|
|
|
|
|
|
|
print(f"[{group_name}] famílias={total_familias} → train(tensors)={m_train}, val(tensors)={m_val}, test(tensors)={m_test}")
|
2026-01-22 18:47:56 +00:00
|
|
|
return {"train": m_train, "val": m_val, "test": m_test, "familias": total_familias}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def main():
|
2026-04-22 20:08:49 +00:00
|
|
|
ap = argparse.ArgumentParser(description="Split estratificado por grupo SEM vazamento (tensors/masks).")
|
|
|
|
|
ap.add_argument("--train", type=float, default=0.70)
|
|
|
|
|
ap.add_argument("--val", type=float, default=0.29)
|
|
|
|
|
ap.add_argument("--test", type=float, default=0.01)
|
|
|
|
|
ap.add_argument("--seed", type=int, default=42)
|
2026-01-22 18:47:56 +00:00
|
|
|
|
2026-04-22 20:08:49 +00:00
|
|
|
ap.add_argument("--min-train", type=int, default=1)
|
|
|
|
|
ap.add_argument("--min-val", type=int, default=1)
|
|
|
|
|
ap.add_argument("--min-test", type=int, default=0)
|
2026-01-22 18:47:56 +00:00
|
|
|
|
2026-04-22 20:08:49 +00:00
|
|
|
ap.add_argument("--resolucao", type=str, default=None,
|
|
|
|
|
help="Sobrescreve resolução no formato WxH (ex: 1024x800).")
|
2026-01-22 18:47:56 +00:00
|
|
|
|
|
|
|
|
ap.add_argument("--cap-train-families", type=str, default="",
|
2026-04-22 20:08:49 +00:00
|
|
|
help="Mapa 'grupo:cap,...' para limitar famílias no TRAIN. Ex.: 'chao:350'")
|
2026-01-22 18:47:56 +00:00
|
|
|
|
|
|
|
|
args = ap.parse_args()
|
|
|
|
|
|
|
|
|
|
if args.resolucao:
|
|
|
|
|
try:
|
|
|
|
|
w, h = args.resolucao.lower().split("x")
|
|
|
|
|
resolucao = (int(w), int(h))
|
|
|
|
|
except Exception:
|
|
|
|
|
resolucao = RESOLUCAO
|
|
|
|
|
else:
|
|
|
|
|
resolucao = RESOLUCAO
|
|
|
|
|
|
|
|
|
|
def parse_cap_map(s):
|
|
|
|
|
caps = {}
|
|
|
|
|
if not s:
|
|
|
|
|
return caps
|
|
|
|
|
for item in s.split(","):
|
|
|
|
|
k, v = item.strip().split(":")
|
|
|
|
|
caps[k.strip()] = int(v)
|
|
|
|
|
return caps
|
|
|
|
|
|
|
|
|
|
caps_map = parse_cap_map(args.cap_train_families)
|
|
|
|
|
|
|
|
|
|
global pasta_origem, pasta_destino
|
2026-04-20 18:53:28 +00:00
|
|
|
pasta_origem = os.path.join("dataset", f"{resolucao[0]}x{resolucao[1]}", "group")
|
|
|
|
|
pasta_destino = os.path.join("dataset", "split")
|
2026-01-22 18:47:56 +00:00
|
|
|
|
|
|
|
|
soma = args.train + args.val + args.test
|
|
|
|
|
if soma <= 0:
|
|
|
|
|
raise ValueError("Soma de proporções deve ser > 0.")
|
2026-04-22 20:08:49 +00:00
|
|
|
|
2026-01-22 18:47:56 +00:00
|
|
|
p_train = args.train / soma
|
2026-04-22 20:08:49 +00:00
|
|
|
p_val = args.val / soma
|
|
|
|
|
p_test = args.test / soma
|
2026-01-22 18:47:56 +00:00
|
|
|
|
|
|
|
|
mins = {
|
|
|
|
|
"train": max(0, args.min_train),
|
2026-04-22 20:08:49 +00:00
|
|
|
"val": max(0, args.min_val),
|
|
|
|
|
"test": max(0, args.min_test),
|
2026-01-22 18:47:56 +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)
|
|
|
|
|
|
|
|
|
|
total_global = {"train": 0, "val": 0, "test": 0, "familias": 0}
|
|
|
|
|
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 (famílias): 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)
|
|
|
|
|
for k in total_global.keys():
|
|
|
|
|
total_global[k] += res.get(k, 0)
|
|
|
|
|
|
2026-04-22 20:08:49 +00:00
|
|
|
print("\nResumo global (tensors copiados):")
|
2026-01-22 18:47:56 +00:00
|
|
|
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']}")
|
2026-04-22 20:08:49 +00:00
|
|
|
print("\n✅ Split sem vazamento concluído!")
|
2026-01-22 18:47:56 +00:00
|
|
|
|
|
|
|
|
|
|
|
|
|
if __name__ == "__main__":
|
2026-04-22 20:08:49 +00:00
|
|
|
main()
|