agrobot_base/Python/OAK/datasets/oak-fcc-3/_7_split.py

3039 lines
70 KiB
Python

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
_7_split_agri_teacher_v2.py
Split agrícola inteligente para OAK-FCC-3 / Teacher V2.
Objetivos:
- VAL/TEST 100% reais.
- Sintéticos/augment/copy-paste somente no TRAIN.
- Máscara semântica é a fonte da verdade para o grupo real.
- Proteção por família para impedir leakage de augmentações.
- Agrupamento conservador de famílias visualmente quase idênticas.
- Estratificação por problema agrícola, não apenas pelo nome da pasta:
* grupo semântico real
* quantidade de erva
* proximidade/contato cana x erva
- Seleção DIVERSA das famílias de VAL/TEST (coreset/farthest-point),
para que a "prova" cubra o espaço de situações em vez de ser um random puro.
- Descritores usam máscara + preview +, opcionalmente, tensor multiespectral
já normalizado/preparado pelo pipeline.
- Gera auditoria completa para sabermos exatamente por que cada família caiu
em TRAIN/VAL/TEST.
A ideia pedagógica:
TRAIN = sala de aula ampla e diversa.
VAL = prova real, representativa e cobrindo os pontos difíceis.
TEST = prova final opcional, também 100% real.
Estrutura esperada:
<root>/<grupo>/tensors/*.npy
<root>/<grupo>/masks/*.npy
<root>/<grupo>/masks_vegetation/*.npy
<root>/<grupo>/masks_cana/*.npy
<root>/<grupo>/metas/*.json
<root>/<grupo>/previews/*.png
<root>/<grupo>/visuals/*.png
Teste
python .\_7_split.py --train 0.85 --val 0.15 --test 0 --strategy agri_diverse --merge-near-duplicate-families --dry-run
python .\_7_split.py --train 0.85 --val 0.15 --test 0 --strategy agri_diverse --merge-near-duplicate-families --use-tensor-features --dry-run
Exemplo recomendado, somente dados reais:
python _7_split.py --train 0.85 --val 0.15 --test 0 --strategy agri_diverse --use-tensor-features --merge-near-duplicate-families --clear-dst
Real + copy/paste:
python _7_split.py ^
--src-roots dataset/1280x800/group,dataset/copypaste/group ^
--train-only-roots dataset/copypaste/group ^
--train 0.85 --val 0.15 --test 0 ^
--strategy agri_diverse ^
--synthetic-train-only ^
--synthetic-respect-family-split ^
--clear-dst
python .\_7_split.py --src-root dataset/960x600/group --train 0.85 --val 0.15 --test 0 --strategy agri_diverse --use-tensor-features --merge-near-duplicate-families --seed 42 --clear-dst
"""
from __future__ import annotations
import argparse
import csv
import json
import math
import os
import random
import re
import shutil
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple
import cv2
import numpy as np
# ============================================================
# Config
# ============================================================
try:
with open("config.json", "r", encoding="utf-8") as f:
CONFIG = json.load(f)
except Exception:
CONFIG = {}
RESOLUCAO = tuple(CONFIG.get("resolucao", [1024, 640]))
MULTI_HEAD = bool(CONFIG.get("multi_head", False))
TENSOR_EXT = ".npy"
MASK_SUFFIX = ".npy"
AUX_MASK_DIRS = [
"masks_vegetation",
"masks_cana",
]
CLASS_ID_TO_NAME = {
0: "chao",
1: "cana",
2: "erva",
}
GROUP_BY_CLASS_SET = {
frozenset({0}): "chao",
frozenset({1}): "cana",
frozenset({2}): "erva",
frozenset({0, 1}): "chao_cana",
frozenset({0, 2}): "chao_erva",
frozenset({1, 2}): "cana_erva",
frozenset({0, 1, 2}): "chao_cana_erva",
}
# ============================================================
# Regex origem/família
# ============================================================
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,
)
RE_COPYPASTE_SUFFIX = re.compile(
r"(.+)_cp_\d+$",
re.IGNORECASE,
)
RE_GROUP_COPYPASTE_SUFFIX = re.compile(
r"(.+)_copypaste$",
re.IGNORECASE,
)
# ============================================================
# Utilidades
# ============================================================
def garantir(path: str | Path) -> None:
os.makedirs(str(path), exist_ok=True)
def limpar_dir(path: str | Path) -> None:
if os.path.isdir(str(path)):
shutil.rmtree(str(path))
garantir(path)
def norm_path(path: str | Path) -> str:
return os.path.normcase(os.path.abspath(str(path)))
def parse_csv_list(value: Optional[str]) -> List[str]:
if not value:
return []
return [x.strip() for x in str(value).split(",") if x.strip()]
def load_json_safe(path: Optional[str | Path]) -> Dict[str, Any]:
if not path:
return {}
try:
with open(path, "r", encoding="utf-8") as f:
data = json.load(f)
return data if isinstance(data, dict) else {}
except Exception:
return {}
def save_json(path: str | Path, data: Dict[str, Any]) -> None:
garantir(os.path.dirname(str(path)))
with open(path, "w", encoding="utf-8") as f:
json.dump(data, f, ensure_ascii=False, indent=2)
def lista_grupos(root: str | Path) -> List[str]:
root = str(root)
if not os.path.isdir(root):
return []
result = []
for group in sorted(os.listdir(root)):
group_dir = os.path.join(root, group)
if not os.path.isdir(group_dir):
continue
if (
os.path.isdir(os.path.join(group_dir, "tensors"))
and os.path.isdir(os.path.join(group_dir, "masks"))
):
result.append(group)
return result
def listar_tensors(tensor_dir: str | Path) -> List[str]:
if not os.path.isdir(str(tensor_dir)):
return []
return sorted(
f
for f in os.listdir(str(tensor_dir))
if f.lower().endswith(TENSOR_EXT)
)
def base_no_ext(filename: str) -> str:
return os.path.splitext(filename)[0]
def normalize_group_for_split(group_name: str) -> str:
m = RE_GROUP_COPYPASTE_SUFFIX.match(group_name)
if m:
return m.group(1)
return group_name
def output_group_name(
source_group: str,
actual_group: str,
is_synthetic: bool,
mode: str,
) -> str:
if mode == "source":
return source_group
if mode == "base":
return actual_group
if mode == "suffix":
return f"{actual_group}_synthetic" if is_synthetic else actual_group
return actual_group
def classify_source_and_family(
filename_no_ext: str,
meta: Optional[Dict[str, Any]] = None,
) -> Tuple[str, str, bool]:
meta = meta or {}
m = RE_ORIGINAL_PREFIX.match(filename_no_ext)
if m:
return "original", m.group(1), False
m = RE_AUGMENTED_FAMILY.match(filename_no_ext)
if m:
return "augmented", m.group(1), True
if RE_AUG_SUFFIX.search(filename_no_ext):
family = RE_AUG_SUFFIX.sub("", filename_no_ext)
return "augmented", family, True
m = RE_COPYPASTE_SUFFIX.match(filename_no_ext)
if m:
return "copypaste", m.group(1), True
cp = meta.get("copy_paste_augmentation")
if isinstance(cp, dict):
receiver = cp.get("receiver")
if isinstance(receiver, dict):
family = str(receiver.get("base") or filename_no_ext)
return "copypaste", family, True
return "copypaste", filename_no_ext, True
if bool(meta.get("synthetic", False)):
return "synthetic", filename_no_ext, True
return "unknown", filename_no_ext, False
def read_mask(path: str | Path) -> np.ndarray:
mask = np.load(str(path), mmap_mode="r")
arr = np.asarray(mask)
if arr.ndim > 2:
arr = np.squeeze(arr)
if arr.ndim != 2:
raise RuntimeError(
f"Máscara semântica deve ser HxW; veio {arr.shape}: {path}"
)
return arr
def infer_actual_group(
mask: np.ndarray,
presence_min_pixels: int,
presence_min_fraction: float,
ignore_ids: Sequence[int],
) -> Tuple[str, Tuple[int, ...]]:
ignore_set = set(int(x) for x in ignore_ids)
unique_ids = set(int(x) for x in np.unique(mask).tolist())
unknown = sorted(
x
for x in unique_ids
if x not in CLASS_ID_TO_NAME and x not in ignore_set
)
if unknown:
raise RuntimeError(
f"Máscara possui IDs desconhecidos: {unknown}"
)
present = []
n = int(mask.size)
for class_id in sorted(CLASS_ID_TO_NAME):
count = int(np.count_nonzero(mask == class_id))
frac = count / max(n, 1)
if (
count >= max(1, int(presence_min_pixels))
and frac >= float(presence_min_fraction)
):
present.append(class_id)
group = GROUP_BY_CLASS_SET.get(frozenset(present))
if group is None:
raise RuntimeError(
f"Combinação de classes não mapeável: {present}"
)
return group, tuple(present)
# ============================================================
# Dataclasses
# ============================================================
@dataclass
class Sample:
src_root: str
source_root_label: str
source_root_train_only: bool
source_group: str
split_group_folder: str
tensor_path: str
tensor_name: str
base: str
mask_path: str
source: str
family: str
is_synthetic: bool
meta_path: Optional[str]
preview_path: Optional[str]
visual_path: Optional[str]
actual_group: str = ""
present_ids: Tuple[int, ...] = field(default_factory=tuple)
frac_chao: float = 0.0
frac_cana: float = 0.0
frac_erva: float = 0.0
frac_ignore: float = 0.0
comp_cana: int = 0
comp_erva: int = 0
largest_cana_frac: float = 0.0
largest_erva_frac: float = 0.0
weed_near_cane_frac: float = 0.0
weed_contact_cane_frac: float = 0.0
brightness: float = 0.0
contrast: float = 0.0
saturation: float = 0.0
dhash64: int = 0
feature_vector: Optional[np.ndarray] = None
stratum: str = ""
agri_risk: float = 0.0
audit_error: str = ""
@dataclass
class FamilyUnit:
actual_group: str
family: str
leakage_id: str
samples: List[Sample]
feature_vector: np.ndarray
stratum: str
agri_risk: float
dhash64: int
frac_chao: float
frac_cana: float
frac_erva: float
assigned_split: Optional[str] = None
assignment_reason: str = ""
# ============================================================
# Coleta
# ============================================================
def collect_samples_from_root(
src_root: str | Path,
train_only_roots_norm: set[str],
label: str,
groups_filter: Optional[set[str]],
) -> List[Sample]:
src_root = str(src_root)
root_norm = norm_path(src_root)
root_train_only = root_norm in train_only_roots_norm
samples: List[Sample] = []
for group_name in lista_grupos(src_root):
split_group_folder = normalize_group_for_split(group_name)
if (
groups_filter
and group_name not in groups_filter
and split_group_folder not in groups_filter
):
continue
group_dir = os.path.join(src_root, group_name)
tensor_dir = os.path.join(group_dir, "tensors")
mask_dir = os.path.join(group_dir, "masks")
meta_dir = os.path.join(group_dir, "metas")
preview_dir = os.path.join(group_dir, "previews")
visual_dir = os.path.join(group_dir, "visuals")
for tensor_name in listar_tensors(tensor_dir):
base = base_no_ext(tensor_name)
tensor_path = os.path.join(tensor_dir, tensor_name)
mask_path = os.path.join(mask_dir, base + MASK_SUFFIX)
if not os.path.isfile(mask_path):
continue
meta_path = os.path.join(meta_dir, base + ".json")
if not os.path.isfile(meta_path):
meta_path = None
meta = load_json_safe(meta_path)
source, family, is_synthetic = classify_source_and_family(
base,
meta,
)
preview_path = None
for ext in (".png", ".jpg", ".jpeg"):
candidate = os.path.join(preview_dir, base + ext)
if os.path.isfile(candidate):
preview_path = candidate
break
visual_path = None
for ext in (".png", ".jpg", ".jpeg"):
candidate = os.path.join(
visual_dir,
base + "_debug" + ext,
)
if os.path.isfile(candidate):
visual_path = candidate
break
candidate = os.path.join(
visual_dir,
base + ext,
)
if os.path.isfile(candidate):
visual_path = candidate
break
samples.append(
Sample(
src_root=src_root,
source_root_label=label,
source_root_train_only=root_train_only,
source_group=group_name,
split_group_folder=split_group_folder,
tensor_path=tensor_path,
tensor_name=tensor_name,
base=base,
mask_path=mask_path,
source=source,
family=family,
is_synthetic=is_synthetic,
meta_path=meta_path,
preview_path=preview_path,
visual_path=visual_path,
)
)
return samples
def collect_all_samples(
src_roots: Sequence[str],
train_only_roots: Sequence[str],
groups: Optional[Sequence[str]],
) -> List[Sample]:
train_only_norm = {
norm_path(x)
for x in train_only_roots
}
groups_filter = set(groups) if groups else None
result = []
for i, root in enumerate(src_roots):
result.extend(
collect_samples_from_root(
src_root=root,
train_only_roots_norm=train_only_norm,
label=f"root{i}",
groups_filter=groups_filter,
)
)
return result
# ============================================================
# Features agrícolas
# ============================================================
def connected_component_stats(binary: np.ndarray) -> Tuple[int, float]:
binary = np.ascontiguousarray(
binary.astype(np.uint8)
)
n, _labels, stats, _centroids = cv2.connectedComponentsWithStats(
binary,
connectivity=8,
)
if n <= 1:
return 0, 0.0
areas = stats[1:, cv2.CC_STAT_AREA].astype(np.float64)
return (
int(len(areas)),
float(np.max(areas) / max(binary.size, 1)),
)
def weed_cane_relation(
mask_small: np.ndarray,
near_radius_frac: float,
) -> Tuple[float, float]:
cane = mask_small == 1
weed = mask_small == 2
if not np.any(cane) or not np.any(weed):
return 0.0, 0.0
h, w = mask_small.shape
radius = max(
1,
int(round(min(h, w) * float(near_radius_frac))),
)
dist = cv2.distanceTransform(
(~cane).astype(np.uint8),
cv2.DIST_L2,
3,
)
weed_dist = dist[weed]
near_frac = (
float(np.mean(weed_dist <= radius))
if weed_dist.size
else 0.0
)
kernel = np.ones((3, 3), np.uint8)
cane_dilated = cv2.dilate(
cane.astype(np.uint8),
kernel,
iterations=1,
) > 0
weed_count = int(np.count_nonzero(weed))
contact = int(np.count_nonzero(weed & cane_dilated))
contact_frac = contact / max(weed_count, 1)
return float(near_frac), float(contact_frac)
def grid_fraction(
binary: np.ndarray,
rows: int = 4,
cols: int = 6,
) -> np.ndarray:
h, w = binary.shape
result = []
for gy in range(rows):
y0 = int(round(gy * h / rows))
y1 = int(round((gy + 1) * h / rows))
for gx in range(cols):
x0 = int(round(gx * w / cols))
x1 = int(round((gx + 1) * w / cols))
tile = binary[y0:y1, x0:x1]
result.append(
float(np.mean(tile))
if tile.size
else 0.0
)
return np.asarray(
result,
dtype=np.float32,
)
def dhash64(img_bgr: np.ndarray) -> int:
gray = cv2.cvtColor(
img_bgr,
cv2.COLOR_BGR2GRAY,
)
small = cv2.resize(
gray,
(9, 8),
interpolation=cv2.INTER_AREA,
)
diff = small[:, 1:] > small[:, :-1]
value = 0
for bit in diff.flatten():
value = (value << 1) | int(bool(bit))
return int(value)
def preview_features(
preview_path: Optional[str],
) -> Tuple[np.ndarray, float, float, float, int]:
if not preview_path:
return (
np.zeros(35, dtype=np.float32),
0.0,
0.0,
0.0,
0,
)
img = cv2.imread(
preview_path,
cv2.IMREAD_COLOR,
)
if img is None:
return (
np.zeros(35, dtype=np.float32),
0.0,
0.0,
0.0,
0,
)
hsv = cv2.cvtColor(
img,
cv2.COLOR_BGR2HSV,
)
brightness = float(
np.mean(hsv[..., 2]) / 255.0
)
saturation = float(
np.mean(hsv[..., 1]) / 255.0
)
gray = cv2.cvtColor(
img,
cv2.COLOR_BGR2GRAY,
)
contrast = float(
np.std(gray) / 255.0
)
hist_h = cv2.calcHist(
[hsv],
[0],
None,
[12],
[0, 180],
).flatten()
hist_s = cv2.calcHist(
[hsv],
[1],
None,
[8],
[0, 256],
).flatten()
hist_v = cv2.calcHist(
[hsv],
[2],
None,
[8],
[0, 256],
).flatten()
hist = np.concatenate([
hist_h,
hist_s,
hist_v,
]).astype(np.float32)
hist_sum = float(hist.sum())
if hist_sum > 0:
hist /= hist_sum
low = cv2.resize(
gray,
(7, 1),
interpolation=cv2.INTER_AREA,
).astype(np.float32).flatten() / 255.0
feat = np.concatenate([
hist,
low,
]).astype(np.float32)
return (
feat,
brightness,
contrast,
saturation,
dhash64(img),
)
def tensor_summary_features(
tensor_path: str,
mask: np.ndarray,
enabled: bool,
stride: int,
) -> np.ndarray:
"""
Features multiespectrais baratas.
Usa mmap + subamostragem espacial.
Para cada canal disponível:
mean, std, p10, p50, p90
Para os 5 primeiros canais físicos, acrescenta:
média na CANA e média na ERVA
Também calcula NDVI/NDRE aproximados a partir de:
R=0, RE=3, NIR=4
quando C>=5.
Não depende de o treino usar 5 ou 7 canais.
"""
if not enabled:
return np.zeros(49, dtype=np.float32)
try:
arr = np.load(
tensor_path,
mmap_mode="r",
)
x = np.asarray(arr)
if x.ndim == 4 and x.shape[0] == 1:
x = x[0]
if x.ndim != 3:
raise RuntimeError(
f"tensor não CHW: {x.shape}"
)
c, h, w = x.shape
step = max(1, int(stride))
ys = np.arange(
0,
h,
step,
dtype=np.int64,
)
xs = np.arange(
0,
w,
step,
dtype=np.int64,
)
x_small = np.asarray(
x[:, ys[:, None], xs[None, :]],
dtype=np.float32,
)
mask_small = cv2.resize(
mask.astype(np.uint8),
(len(xs), len(ys)),
interpolation=cv2.INTER_NEAREST,
)
feats = []
max_channels = min(c, 7)
for ch in range(7):
if ch < max_channels:
values = x_small[ch].reshape(-1)
values = values[np.isfinite(values)]
if values.size:
feats.extend([
float(np.mean(values)),
float(np.std(values)),
float(np.percentile(values, 10)),
float(np.percentile(values, 50)),
float(np.percentile(values, 90)),
])
else:
feats.extend([0.0] * 5)
else:
feats.extend([0.0] * 5)
# médias por classe dos 5 canais físicos
for class_id in (1, 2):
class_mask = mask_small == class_id
for ch in range(5):
if ch < c and np.any(class_mask):
vals = x_small[ch][class_mask]
vals = vals[np.isfinite(vals)]
feats.append(
float(np.mean(vals))
if vals.size
else 0.0
)
else:
feats.append(0.0)
# NDVI / NDRE globais do Raw5 físico.
if c >= 5:
r = x_small[0]
re = x_small[3]
nir = x_small[4]
def nd(a, b):
den = a + b
out = np.zeros_like(a, dtype=np.float32)
np.divide(
a - b,
den,
out=out,
where=np.abs(den) > 1e-6,
)
return np.nan_to_num(
out,
nan=0.0,
posinf=1.0,
neginf=-1.0,
)
ndvi = nd(nir, r)
ndre = nd(nir, re)
feats.extend([
float(np.mean(ndvi)),
float(np.std(ndvi)),
float(np.mean(ndre)),
float(np.std(ndre)),
])
else:
feats.extend([0.0] * 4)
# 35 + 10 + 4 = 49
return np.asarray(
feats,
dtype=np.float32,
)
except Exception:
return np.zeros(
49,
dtype=np.float32,
)
def target_size_bin(
frac_erva: float,
tiny: float,
small: float,
medium: float,
) -> str:
if frac_erva <= 0.0:
return "none"
if frac_erva < tiny:
return "tiny"
if frac_erva < small:
return "small"
if frac_erva < medium:
return "medium"
return "large"
def compute_agri_risk(
frac_cana: float,
frac_erva: float,
weed_near_cane_frac: float,
weed_contact_cane_frac: float,
) -> float:
"""
Risco agrícola relativo.
Não é score de qualidade.
É somente uma forma de marcar casos que são "prova difícil":
- cana + erva simultâneas;
- erva perto/encostada na cana;
- target pequeno em presença de cana.
"""
has_cana = frac_cana > 0.0
has_erva = frac_erva > 0.0
score = 0.0
if has_cana and has_erva:
score += 0.35
score += 0.30 * float(weed_near_cane_frac)
score += 0.20 * float(weed_contact_cane_frac)
if has_cana and has_erva and frac_erva < 0.02:
score += 0.15
return float(
max(0.0, min(score, 1.0))
)
def analyze_sample(
sample: Sample,
args,
) -> None:
mask = read_mask(
sample.mask_path
)
actual_group, present_ids = infer_actual_group(
mask=mask,
presence_min_pixels=args.presence_min_pixels,
presence_min_fraction=args.presence_min_fraction,
ignore_ids=args.ignore_ids,
)
sample.actual_group = actual_group
sample.present_ids = present_ids
sample.frac_chao = float(
np.mean(mask == 0)
)
sample.frac_cana = float(
np.mean(mask == 1)
)
sample.frac_erva = float(
np.mean(mask == 2)
)
sample.frac_ignore = float(
sum(
np.mean(mask == int(i))
for i in args.ignore_ids
)
)
# Para morfologia, reduz bastante. A feature é relativa, não pixel-exata.
small_w = min(320, mask.shape[1])
small_h = max(
1,
int(round(
mask.shape[0]
* small_w
/ max(mask.shape[1], 1)
)),
)
mask_small = cv2.resize(
mask.astype(np.uint8),
(small_w, small_h),
interpolation=cv2.INTER_NEAREST,
)
(
sample.comp_cana,
sample.largest_cana_frac,
) = connected_component_stats(
mask_small == 1
)
(
sample.comp_erva,
sample.largest_erva_frac,
) = connected_component_stats(
mask_small == 2
)
(
sample.weed_near_cane_frac,
sample.weed_contact_cane_frac,
) = weed_cane_relation(
mask_small,
args.near_radius_frac,
)
(
visual_feat,
sample.brightness,
sample.contrast,
sample.saturation,
sample.dhash64,
) = preview_features(
sample.preview_path
)
tensor_feat = tensor_summary_features(
tensor_path=sample.tensor_path,
mask=mask,
enabled=bool(args.use_tensor_features),
stride=args.tensor_feature_stride,
)
cana_grid = grid_fraction(
mask_small == 1,
)
erva_grid = grid_fraction(
mask_small == 2,
)
size_bin = target_size_bin(
sample.frac_erva,
args.target_tiny_frac,
args.target_small_frac,
args.target_medium_frac,
)
if (
sample.frac_cana > 0
and sample.frac_erva > 0
and sample.weed_near_cane_frac >= args.close_weed_cane_threshold
):
proximity_bin = "close"
elif sample.frac_cana > 0 and sample.frac_erva > 0:
proximity_bin = "separate"
else:
proximity_bin = "na"
sample.stratum = (
f"{sample.actual_group}"
f"|target={size_bin}"
f"|caneweed={proximity_bin}"
)
sample.agri_risk = compute_agri_risk(
frac_cana=sample.frac_cana,
frac_erva=sample.frac_erva,
weed_near_cane_frac=sample.weed_near_cane_frac,
weed_contact_cane_frac=sample.weed_contact_cane_frac,
)
scalar_feat = np.asarray([
sample.frac_chao,
sample.frac_cana,
sample.frac_erva,
sample.frac_ignore,
math.log1p(sample.comp_cana) / 6.0,
math.log1p(sample.comp_erva) / 6.0,
sample.largest_cana_frac,
sample.largest_erva_frac,
sample.weed_near_cane_frac,
sample.weed_contact_cane_frac,
sample.brightness,
sample.contrast,
sample.saturation,
sample.agri_risk,
], dtype=np.float32)
sample.feature_vector = np.concatenate([
scalar_feat,
cana_grid,
erva_grid,
visual_feat,
tensor_feat,
]).astype(np.float32)
# ============================================================
# Família / near-duplicate leakage
# ============================================================
def hamming64(a: int, b: int) -> int:
return int(
(int(a) ^ int(b)).bit_count()
)
class UnionFind:
def __init__(self, n: int):
self.parent = list(range(n))
self.rank = [0] * n
def find(self, x: int) -> int:
while self.parent[x] != x:
self.parent[x] = self.parent[
self.parent[x]
]
x = self.parent[x]
return x
def union(self, a: int, b: int) -> None:
ra = self.find(a)
rb = self.find(b)
if ra == rb:
return
if self.rank[ra] < self.rank[rb]:
ra, rb = rb, ra
self.parent[rb] = ra
if self.rank[ra] == self.rank[rb]:
self.rank[ra] += 1
def build_real_family_units(
samples: Sequence[Sample],
) -> List[FamilyUnit]:
grouped: Dict[Tuple[str, str], List[Sample]] = {}
for s in samples:
if s.is_synthetic:
continue
if s.source_root_train_only:
continue
grouped.setdefault(
(s.actual_group, s.family),
[],
).append(s)
units = []
for (group, family), members in grouped.items():
vectors = np.stack([
s.feature_vector
for s in members
], axis=0)
representative = max(
members,
key=lambda s: s.agri_risk,
)
stratum_counts: Dict[str, int] = {}
for s in members:
stratum_counts[s.stratum] = (
stratum_counts.get(s.stratum, 0) + 1
)
stratum = max(
stratum_counts.items(),
key=lambda kv: (kv[1], kv[0]),
)[0]
units.append(
FamilyUnit(
actual_group=group,
family=family,
leakage_id=family,
samples=list(members),
feature_vector=np.mean(
vectors,
axis=0,
).astype(np.float32),
stratum=stratum,
agri_risk=float(
max(s.agri_risk for s in members)
),
dhash64=representative.dhash64,
frac_chao=float(
np.mean([s.frac_chao for s in members])
),
frac_cana=float(
np.mean([s.frac_cana for s in members])
),
frac_erva=float(
np.mean([s.frac_erva for s in members])
),
)
)
return units
def merge_near_duplicate_family_units(
units: Sequence[FamilyUnit],
enabled: bool,
max_hamming: int,
mask_l1: float,
) -> Tuple[List[FamilyUnit], Dict[Tuple[str, str], str], int]:
"""
Une famílias reais muito semelhantes para que não possam cair uma no
TRAIN e outra no VAL.
É propositalmente conservador.
Só compara famílias do mesmo grupo real.
"""
units = list(units)
if not enabled or len(units) <= 1:
mapping = {
(u.actual_group, u.family): u.leakage_id
for u in units
}
return units, mapping, 0
result_units = []
family_to_leakage: Dict[Tuple[str, str], str] = {}
merged_count = 0
by_group: Dict[str, List[FamilyUnit]] = {}
for u in units:
by_group.setdefault(
u.actual_group,
[],
).append(u)
for group, group_units in by_group.items():
n = len(group_units)
uf = UnionFind(n)
# bucket conservador por frações de classes para evitar O(n²) total
buckets: Dict[Tuple[int, int, int], List[int]] = {}
for i, u in enumerate(group_units):
key = (
int(round(u.frac_chao / 0.02)),
int(round(u.frac_cana / 0.02)),
int(round(u.frac_erva / 0.02)),
)
buckets.setdefault(key, []).append(i)
for key, indices in buckets.items():
# compara bucket e vizinhos imediatos de composição
candidate_indices = []
a0, a1, a2 = key
for d0 in (-1, 0, 1):
for d1 in (-1, 0, 1):
for d2 in (-1, 0, 1):
candidate_indices.extend(
buckets.get(
(a0 + d0, a1 + d1, a2 + d2),
[],
)
)
candidate_indices = sorted(
set(candidate_indices)
)
for i in indices:
ui = group_units[i]
for j in candidate_indices:
if j <= i:
continue
uj = group_units[j]
if (
hamming64(
ui.dhash64,
uj.dhash64,
)
> int(max_hamming)
):
continue
l1 = (
abs(ui.frac_chao - uj.frac_chao)
+ abs(ui.frac_cana - uj.frac_cana)
+ abs(ui.frac_erva - uj.frac_erva)
)
if l1 > float(mask_l1):
continue
uf.union(i, j)
clusters: Dict[int, List[FamilyUnit]] = {}
for i, u in enumerate(group_units):
clusters.setdefault(
uf.find(i),
[],
).append(u)
for cluster_index, members in enumerate(clusters.values()):
if len(members) > 1:
merged_count += len(members) - 1
leakage_id = (
members[0].family
if len(members) == 1
else f"near_{group}_{cluster_index:05d}"
)
all_samples = []
for m in members:
all_samples.extend(
m.samples
)
family_to_leakage[
(group, m.family)
] = leakage_id
vectors = np.stack([
m.feature_vector
for m in members
], axis=0)
hardest = max(
members,
key=lambda m: m.agri_risk,
)
stratum_counts: Dict[str, int] = {}
for m in members:
stratum_counts[m.stratum] = (
stratum_counts.get(m.stratum, 0) + 1
)
stratum = max(
stratum_counts.items(),
key=lambda kv: (kv[1], kv[0]),
)[0]
result_units.append(
FamilyUnit(
actual_group=group,
family=members[0].family,
leakage_id=leakage_id,
samples=all_samples,
feature_vector=np.mean(
vectors,
axis=0,
).astype(np.float32),
stratum=stratum,
agri_risk=float(
max(m.agri_risk for m in members)
),
dhash64=hardest.dhash64,
frac_chao=float(
np.mean([m.frac_chao for m in members])
),
frac_cana=float(
np.mean([m.frac_cana for m in members])
),
frac_erva=float(
np.mean([m.frac_erva for m in members])
),
)
)
return (
result_units,
family_to_leakage,
merged_count,
)
# ============================================================
# Seleção diversa de VAL/TEST
# ============================================================
def robust_standardize(
features: np.ndarray,
) -> np.ndarray:
x = np.asarray(
features,
dtype=np.float32,
)
med = np.median(
x,
axis=0,
)
q25 = np.percentile(
x,
25,
axis=0,
)
q75 = np.percentile(
x,
75,
axis=0,
)
scale = q75 - q25
std = np.std(
x,
axis=0,
)
scale = np.where(
scale > 1e-6,
scale,
std,
)
scale = np.where(
scale > 1e-6,
scale,
1.0,
)
z = (x - med) / scale
return np.clip(
z,
-8.0,
8.0,
).astype(np.float32)
def diverse_order(
units: Sequence[FamilyUnit],
seed: int,
) -> List[int]:
"""
Ordem k-center:
começa pela unidade mais próxima do centro do estrato
e depois vai pegando a mais distante do conjunto já coberto.
Isso produz uma prova representativa + ampla, não somente casos extremos.
"""
units = list(units)
if len(units) <= 1:
return list(range(len(units)))
x = np.stack([
u.feature_vector
for u in units
], axis=0)
z = robust_standardize(x)
centroid = np.mean(
z,
axis=0,
)
dist_center = np.sum(
(z - centroid) ** 2,
axis=1,
)
rng = np.random.default_rng(
int(seed)
)
jitter = rng.uniform(
0.0,
1e-8,
size=len(units),
)
first = int(
np.argmin(
dist_center + jitter
)
)
order = [first]
selected = np.zeros(
len(units),
dtype=bool,
)
selected[first] = True
delta = z - z[first]
min_dist_sq = np.sum(
delta * delta,
axis=1,
)
min_dist_sq[first] = 0.0
while len(order) < len(units):
score = min_dist_sq.copy()
score[selected] = -1.0
idx = int(
np.argmax(
score + jitter
)
)
order.append(idx)
selected[idx] = True
delta = z - z[idx]
dist_sq = np.sum(
delta * delta,
axis=1,
)
min_dist_sq = np.minimum(
min_dist_sq,
dist_sq,
)
return order
def allocate_stratum_counts(
n: int,
p_train: float,
p_val: float,
p_test: float,
) -> Tuple[int, int, int]:
if n <= 0:
return 0, 0, 0
n_val = int(round(
n * p_val
))
n_test = int(round(
n * p_test
))
# Garante pelo menos uma família de val em estratos minimamente grandes.
if p_val > 0 and n >= 5:
n_val = max(
1,
n_val,
)
if p_test > 0 and n >= 10:
n_test = max(
1,
n_test,
)
while n_val + n_test >= n:
if n_test > 0:
n_test -= 1
elif n_val > 0:
n_val -= 1
else:
break
n_train = n - n_val - n_test
return (
n_train,
n_val,
n_test,
)
def assign_agri_diverse_split(
units: Sequence[FamilyUnit],
p_train: float,
p_val: float,
p_test: float,
seed: int,
) -> Dict[Tuple[str, str], str]:
"""
Estratifica por "situação agrícola" e escolhe VAL/TEST por diversidade.
Dentro de cada estrato:
- escolhe representantes diversos para VAL;
- escolhe representantes ainda não usados para TEST;
- restante vai para TRAIN.
"""
strata: Dict[str, List[FamilyUnit]] = {}
for u in units:
strata.setdefault(
u.stratum,
[],
).append(u)
assignment: Dict[Tuple[str, str], str] = {}
for stratum_index, stratum in enumerate(sorted(strata)):
members = strata[stratum]
n_train, n_val, n_test = allocate_stratum_counts(
len(members),
p_train,
p_val,
p_test,
)
order = diverse_order(
members,
seed=seed + stratum_index * 9973,
)
# Distribui as posições mais representativas/diversas para a prova.
val_idx = set(
order[:n_val]
)
test_idx = set(
order[n_val:n_val + n_test]
)
for i, unit in enumerate(members):
if i in val_idx:
split = "val"
reason = "agri_diverse_val"
elif i in test_idx:
split = "test"
reason = "agri_diverse_test"
else:
split = "train"
reason = "agri_diverse_train"
unit.assigned_split = split
unit.assignment_reason = reason
assignment[
(unit.actual_group, unit.leakage_id)
] = split
return assignment
def assign_random_split(
units: Sequence[FamilyUnit],
p_train: float,
p_val: float,
p_test: float,
seed: int,
) -> Dict[Tuple[str, str], str]:
by_group: Dict[str, List[FamilyUnit]] = {}
for u in units:
by_group.setdefault(
u.actual_group,
[],
).append(u)
assignment = {}
for group_index, group in enumerate(sorted(by_group)):
members = list(
by_group[group]
)
rng = random.Random(
seed + group_index * 1009
)
rng.shuffle(members)
n = len(members)
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:
n_test = 0
n_train = max(
0,
n - n_val,
)
for i, unit in enumerate(members):
if i < n_train:
split = "train"
elif i < n_train + n_val:
split = "val"
else:
split = "test"
unit.assigned_split = split
unit.assignment_reason = "random_group"
assignment[
(unit.actual_group, unit.leakage_id)
] = split
return assignment
# ============================================================
# Aplicar split aos samples
# ============================================================
def decide_sample_split(
sample: Sample,
family_to_leakage: Dict[Tuple[str, str], str],
assignment: Dict[Tuple[str, str], str],
args,
) -> Tuple[Optional[str], str]:
# Para dados reais, a máscara define o grupo/família de leakage.
# Para sintéticos/copy-paste, a família deve respeitar o split da
# FAMÍLIA REAL DE ORIGEM/receiver. O grupo sintético pode mudar depois
# do copy-paste (ex.: chao_cana -> chao_cana_erva), então usar
# actual_group aqui faria o filho perder o vínculo com o pai real.
family_lookup_group = (
sample.actual_group
if not sample.is_synthetic
else normalize_group_for_split(sample.source_group)
)
leakage_id = family_to_leakage.get(
(family_lookup_group, sample.family),
sample.family,
)
assigned = assignment.get(
(family_lookup_group, leakage_id)
)
is_train_only_candidate = (
sample.source_root_train_only
or (
bool(args.synthetic_train_only)
and sample.is_synthetic
)
or sample.source in (
"augmented",
"copypaste",
"synthetic",
)
)
if is_train_only_candidate:
if bool(args.synthetic_respect_family_split):
if assigned is None:
if bool(args.allow_orphan_synthetic_train):
return (
"train",
"train_only_orphan_allowed",
)
return (
None,
"skip_train_only_orphan_no_real_family",
)
if assigned != "train":
return (
None,
f"skip_train_only_family_assigned_{assigned}",
)
return (
"train",
"train_only_family_train",
)
return (
"train",
"train_only_forced",
)
if assigned is None:
return (
None,
"skip_real_no_family_assignment",
)
return (
assigned,
f"real_assigned_{assigned}",
)
# ============================================================
# Cópia
# ============================================================
def copiar_sample(
sample: Sample,
dst_root: str,
split_name: str,
out_group_name: str,
copy_meta_preview: bool,
copy_visuals: bool,
) -> Optional[Dict[str, Any]]:
dst_group_dir = os.path.join(
dst_root,
split_name,
"group",
out_group_name,
)
dst_tensor_dir = os.path.join(
dst_group_dir,
"tensors",
)
dst_mask_dir = os.path.join(
dst_group_dir,
"masks",
)
garantir(
dst_tensor_dir
)
garantir(
dst_mask_dir
)
if not (
os.path.isfile(sample.tensor_path)
and os.path.isfile(sample.mask_path)
):
return None
tensor_dst = os.path.join(
dst_tensor_dir,
sample.tensor_name,
)
mask_dst = os.path.join(
dst_mask_dir,
os.path.basename(sample.mask_path),
)
shutil.copy2(
sample.tensor_path,
tensor_dst,
)
shutil.copy2(
sample.mask_path,
mask_dst,
)
base = sample.base
mask_png_src = os.path.join(
os.path.dirname(sample.mask_path),
base + ".png",
)
mask_png_dst = None
if os.path.isfile(mask_png_src):
mask_png_dst = os.path.join(
dst_mask_dir,
base + ".png",
)
shutil.copy2(
mask_png_src,
mask_png_dst,
)
aux_masks: Dict[str, Dict[str, Optional[str]]] = {}
src_group_dir = os.path.join(
sample.src_root,
sample.source_group,
)
for aux_dir in AUX_MASK_DIRS:
aux_src_dir = os.path.join(
src_group_dir,
aux_dir,
)
aux_dst_dir = os.path.join(
dst_group_dir,
aux_dir,
)
aux_npy_src = os.path.join(
aux_src_dir,
base + ".npy",
)
aux_png_src = os.path.join(
aux_src_dir,
base + ".png",
)
aux_npy_dst = None
aux_png_dst = None
if os.path.isfile(aux_npy_src):
garantir(
aux_dst_dir
)
aux_npy_dst = os.path.join(
aux_dst_dir,
base + ".npy",
)
shutil.copy2(
aux_npy_src,
aux_npy_dst,
)
if os.path.isfile(aux_png_src):
aux_png_dst = os.path.join(
aux_dst_dir,
base + ".png",
)
shutil.copy2(
aux_png_src,
aux_png_dst,
)
aux_masks[aux_dir] = {
"npy": aux_npy_dst,
"png": aux_png_dst,
}
if MULTI_HEAD and aux_npy_dst is None:
raise RuntimeError(
f"multi_head=true e máscara auxiliar ausente: "
f"{aux_dir}/{base}.npy"
)
meta_dst = None
preview_dst = None
visual_dst = None
if copy_meta_preview:
if (
sample.meta_path
and os.path.isfile(sample.meta_path)
):
dst_meta_dir = os.path.join(
dst_group_dir,
"metas",
)
garantir(
dst_meta_dir
)
meta_dst = os.path.join(
dst_meta_dir,
base + ".json",
)
shutil.copy2(
sample.meta_path,
meta_dst,
)
if (
sample.preview_path
and os.path.isfile(sample.preview_path)
):
dst_preview_dir = os.path.join(
dst_group_dir,
"previews",
)
garantir(
dst_preview_dir
)
ext = os.path.splitext(
sample.preview_path
)[1]
preview_dst = os.path.join(
dst_preview_dir,
base + ext,
)
shutil.copy2(
sample.preview_path,
preview_dst,
)
if (
copy_visuals
and sample.visual_path
and os.path.isfile(sample.visual_path)
):
dst_visual_dir = os.path.join(
dst_group_dir,
"visuals",
)
garantir(
dst_visual_dir
)
visual_dst = os.path.join(
dst_visual_dir,
os.path.basename(sample.visual_path),
)
shutil.copy2(
sample.visual_path,
visual_dst,
)
return {
"split": split_name,
"group": out_group_name,
"source_group": sample.source_group,
"actual_group": sample.actual_group,
"base": base,
"family": sample.family,
"source": sample.source,
"is_synthetic": int(sample.is_synthetic),
"stratum": sample.stratum,
"agri_risk": sample.agri_risk,
"frac_chao": sample.frac_chao,
"frac_cana": sample.frac_cana,
"frac_erva": sample.frac_erva,
"weed_near_cane_frac": sample.weed_near_cane_frac,
"weed_contact_cane_frac": sample.weed_contact_cane_frac,
"source_root": sample.src_root,
"tensor": tensor_dst,
"mask_npy": mask_dst,
"mask_png": mask_png_dst,
"mask_vegetation_npy": aux_masks.get(
"masks_vegetation",
{},
).get("npy"),
"mask_cana_npy": aux_masks.get(
"masks_cana",
{},
).get("npy"),
"meta": meta_dst,
"preview": preview_dst,
"visual_debug": visual_dst,
}
# ============================================================
# Relatórios
# ============================================================
def write_csv(
path: str | Path,
rows: Sequence[Dict[str, Any]],
fieldnames: Sequence[str],
) -> None:
garantir(
os.path.dirname(str(path))
)
with open(
path,
"w",
newline="",
encoding="utf-8",
) as f:
writer = csv.DictWriter(
f,
fieldnames=list(fieldnames),
extrasaction="ignore",
)
writer.writeheader()
writer.writerows(rows)
def sample_audit_row(
sample: Sample,
split_name: Optional[str],
reason: str,
) -> Dict[str, Any]:
return {
"source_group": sample.source_group,
"actual_group": sample.actual_group,
"folder_mask_mismatch": (
normalize_group_for_split(sample.source_group)
!= sample.actual_group
),
"base": sample.base,
"family": sample.family,
"source": sample.source,
"is_synthetic": int(sample.is_synthetic),
"split": split_name or "",
"reason": reason,
"stratum": sample.stratum,
"agri_risk": sample.agri_risk,
"present_ids": ",".join(
map(str, sample.present_ids)
),
"frac_chao": sample.frac_chao,
"frac_cana": sample.frac_cana,
"frac_erva": sample.frac_erva,
"frac_ignore": sample.frac_ignore,
"comp_cana": sample.comp_cana,
"comp_erva": sample.comp_erva,
"largest_cana_frac": sample.largest_cana_frac,
"largest_erva_frac": sample.largest_erva_frac,
"weed_near_cane_frac": sample.weed_near_cane_frac,
"weed_contact_cane_frac": sample.weed_contact_cane_frac,
"brightness": sample.brightness,
"contrast": sample.contrast,
"saturation": sample.saturation,
"dhash64": f"{sample.dhash64:016x}",
"tensor": sample.tensor_path,
"mask": sample.mask_path,
"preview": sample.preview_path or "",
"audit_error": sample.audit_error,
}
def family_audit_rows(
units: Sequence[FamilyUnit],
) -> List[Dict[str, Any]]:
rows = []
for u in sorted(
units,
key=lambda x: (
x.actual_group,
x.stratum,
x.leakage_id,
),
):
rows.append({
"actual_group": u.actual_group,
"family": u.family,
"leakage_id": u.leakage_id,
"samples_in_unit": len(u.samples),
"stratum": u.stratum,
"agri_risk": u.agri_risk,
"frac_chao": u.frac_chao,
"frac_cana": u.frac_cana,
"frac_erva": u.frac_erva,
"assigned_split": u.assigned_split or "",
"assignment_reason": u.assignment_reason,
})
return rows
# ============================================================
# Main
# ============================================================
def main() -> None:
ap = argparse.ArgumentParser(
description=(
"Split agrícola Teacher V2: família, máscara, diversidade e "
"prova real sem leakage."
)
)
ap.add_argument(
"--train",
type=float,
default=0.85,
)
ap.add_argument(
"--val",
type=float,
default=0.15,
)
ap.add_argument(
"--test",
type=float,
default=0.0,
)
ap.add_argument(
"--seed",
type=int,
default=42,
)
ap.add_argument(
"--strategy",
choices=[
"agri_diverse",
"random",
],
default="agri_diverse",
)
ap.add_argument(
"--resolucao",
default=None,
help="WxH, ex. 1280x800",
)
ap.add_argument(
"--src-root",
default=None,
)
ap.add_argument(
"--src-roots",
default="",
)
ap.add_argument(
"--train-only-roots",
default="",
)
ap.add_argument(
"--dst-root",
default="dataset/split",
)
ap.add_argument(
"--groups",
default=None,
)
ap.add_argument(
"--clear-dst",
action="store_true",
)
ap.add_argument(
"--no-meta-preview",
action="store_true",
)
ap.add_argument(
"--no-visuals",
action="store_true",
)
ap.add_argument(
"--output-group-mode",
choices=[
"source",
"base",
"suffix",
],
default="base",
help=(
"base é recomendado: o grupo final segue a máscara real."
),
)
# máscara
ap.add_argument(
"--presence-min-pixels",
type=int,
default=1,
)
ap.add_argument(
"--presence-min-fraction",
type=float,
default=0.0,
)
ap.add_argument(
"--ignore-ids",
nargs="*",
type=int,
default=[255],
)
# Teacher/agri strata
ap.add_argument(
"--target-tiny-frac",
type=float,
default=0.002,
)
ap.add_argument(
"--target-small-frac",
type=float,
default=0.010,
)
ap.add_argument(
"--target-medium-frac",
type=float,
default=0.050,
)
ap.add_argument(
"--near-radius-frac",
type=float,
default=0.02,
)
ap.add_argument(
"--close-weed-cane-threshold",
type=float,
default=0.25,
)
# features
ap.add_argument(
"--use-tensor-features",
action="store_true",
default=False,
help=(
"Inclui diversidade espectral do tensor no split. "
"Recomendado após validar custo no dry-run."
),
)
ap.add_argument(
"--tensor-feature-stride",
type=int,
default=16,
help=(
"Subamostragem espacial para features do tensor. "
"16 é leve mesmo em 1280x800."
),
)
# near duplicates
ap.add_argument(
"--merge-near-duplicate-families",
action="store_true",
default=False,
help=(
"Une famílias reais quase idênticas para impedir leakage."
),
)
ap.add_argument(
"--near-dup-hamming",
type=int,
default=2,
)
ap.add_argument(
"--near-dup-mask-l1",
type=float,
default=0.02,
)
# synthetic
ap.add_argument(
"--synthetic-train-only",
action="store_true",
default=True,
)
ap.add_argument(
"--allow-synthetic-val",
dest="synthetic_train_only",
action="store_false",
)
ap.add_argument(
"--synthetic-respect-family-split",
action="store_true",
default=True,
)
ap.add_argument(
"--no-synthetic-respect-family-split",
dest="synthetic_respect_family_split",
action="store_false",
)
ap.add_argument(
"--allow-orphan-synthetic-train",
action="store_true",
default=False,
)
# relatórios
ap.add_argument(
"--manifest",
default="",
)
ap.add_argument(
"--audit",
default="",
)
ap.add_argument(
"--families",
default="",
)
ap.add_argument(
"--summary",
default="",
)
ap.add_argument(
"--skipped",
default="",
)
ap.add_argument(
"--dry-run",
action="store_true",
help="Audita e decide split, mas não copia arquivos.",
)
args = ap.parse_args()
# resolução
if args.resolucao:
try:
w, h = args.resolucao.lower().split("x")
resolution = (
int(w),
int(h),
)
except Exception:
raise SystemExit(
f"[ERRO] resolução inválida: {args.resolucao}"
)
else:
resolution = RESOLUCAO
default_src = os.path.join(
"dataset",
f"{resolution[0]}x{resolution[1]}",
"group",
)
src_roots = parse_csv_list(
args.src_roots
)
if args.src_root:
src_roots.insert(
0,
args.src_root,
)
if not src_roots:
src_roots = [
default_src
]
seen = set()
unique_src_roots = []
for root in src_roots:
nr = norm_path(root)
if nr not in seen:
seen.add(nr)
unique_src_roots.append(root)
src_roots = unique_src_roots
train_only_roots = parse_csv_list(
args.train_only_roots
)
for root in src_roots:
if not os.path.isdir(root):
raise SystemExit(
f"[ERRO] src-root não encontrado: {root}"
)
soma = (
args.train
+ args.val
+ args.test
)
if soma <= 0:
raise SystemExit(
"[ERRO] soma train+val+test deve ser > 0"
)
p_train = args.train / soma
p_val = args.val / soma
p_test = args.test / soma
groups = (
parse_csv_list(args.groups)
if args.groups
else None
)
dst_root = args.dst_root
if (
args.clear_dst
and not args.dry_run
):
print(
f"[INFO] Limpando destino: {dst_root}"
)
limpar_dir(
dst_root
)
else:
garantir(
dst_root
)
print("==========================================")
print("SPLIT AGRÍCOLA TEACHER V2")
print(f"Resolution : {resolution[0]}x{resolution[1]}")
print(f"Strategy : {args.strategy}")
print(
f"Split : train={p_train:.3f} "
f"val={p_val:.3f} test={p_test:.3f}"
)
print(f"Tensor feat: {args.use_tensor_features}")
print(
f"Near-dups : {args.merge_near_duplicate_families}"
)
print(f"Dry-run : {args.dry_run}")
print("SRC:")
for root in src_roots:
marker = (
" [TRAIN_ONLY]"
if norm_path(root)
in {norm_path(x) for x in train_only_roots}
else ""
)
print(f" - {root}{marker}")
print("==========================================")
print("[1/6] Coletando samples...")
samples = collect_all_samples(
src_roots=src_roots,
train_only_roots=train_only_roots,
groups=groups,
)
if not samples:
raise SystemExit(
"[ERRO] Nenhum sample encontrado."
)
print(
f"[OK] samples encontrados: {len(samples)}"
)
print("[2/6] Auditando máscaras + features agrícolas...")
analyzed: List[Sample] = []
rejected: List[Sample] = []
for i, sample in enumerate(
samples,
start=1,
):
try:
analyze_sample(
sample,
args,
)
analyzed.append(
sample
)
except Exception as exc:
sample.audit_error = str(exc)
rejected.append(
sample
)
if (
i % 100 == 0
or i == len(samples)
):
print(
f" {i}/{len(samples)} "
f"ok={len(analyzed)} "
f"rejeitados={len(rejected)}"
)
real_analyzed = [
s
for s in analyzed
if not s.is_synthetic
and not s.source_root_train_only
]
mismatches = [
s
for s in real_analyzed
if normalize_group_for_split(
s.source_group
) != s.actual_group
]
print(
f"[OK] reais auditados={len(real_analyzed)} "
f"mismatch pasta/máscara={len(mismatches)}"
)
print("[3/6] Construindo famílias reais...")
base_units = build_real_family_units(
real_analyzed
)
(
units,
family_to_leakage,
merged_count,
) = merge_near_duplicate_family_units(
units=base_units,
enabled=args.merge_near_duplicate_families,
max_hamming=args.near_dup_hamming,
mask_l1=args.near_dup_mask_l1,
)
# Garante mapping mesmo sem merge.
for unit in units:
for sample in unit.samples:
family_to_leakage.setdefault(
(sample.actual_group, sample.family),
unit.leakage_id,
)
print(
f"[OK] famílias base={len(base_units)} "
f"unidades leakage={len(units)} "
f"famílias mescladas={merged_count}"
)
print("[4/6] Montando a prova...")
if args.strategy == "agri_diverse":
assignment = assign_agri_diverse_split(
units=units,
p_train=p_train,
p_val=p_val,
p_test=p_test,
seed=args.seed,
)
else:
assignment = assign_random_split(
units=units,
p_train=p_train,
p_val=p_val,
p_test=p_test,
seed=args.seed,
)
family_counts = {
"train": 0,
"val": 0,
"test": 0,
}
for unit in units:
if unit.assigned_split:
family_counts[
unit.assigned_split
] += 1
print(
f"[OK] famílias/unidades → "
f"train={family_counts['train']} "
f"val={family_counts['val']} "
f"test={family_counts['test']}"
)
print("[5/6] Aplicando split aos samples...")
decisions: Dict[int, Tuple[Optional[str], str]] = {}
totals = {
"train": 0,
"val": 0,
"test": 0,
"skipped": 0,
"synthetic_train": 0,
"real_train": 0,
"real_val": 0,
"real_test": 0,
}
stratum_counts: Dict[str, Dict[str, int]] = {}
for sample in analyzed:
split_name, reason = decide_sample_split(
sample=sample,
family_to_leakage=family_to_leakage,
assignment=assignment,
args=args,
)
decisions[id(sample)] = (
split_name,
reason,
)
if split_name is None:
totals["skipped"] += 1
continue
totals[split_name] += 1
if (
split_name == "train"
and sample.is_synthetic
):
totals["synthetic_train"] += 1
elif split_name == "train":
totals["real_train"] += 1
elif split_name == "val":
totals["real_val"] += 1
elif split_name == "test":
totals["real_test"] += 1
ss = stratum_counts.setdefault(
sample.stratum,
{
"train": 0,
"val": 0,
"test": 0,
},
)
ss[split_name] += 1
print(
f"[OK] samples → "
f"train={totals['train']} "
f"val={totals['val']} "
f"test={totals['test']} "
f"skipped={totals['skipped']}"
)
print("[6/6] Copiando / relatórios...")
manifest_rows = []
if not args.dry_run:
for sample in analyzed:
split_name, reason = decisions[
id(sample)
]
if split_name is None:
continue
out_group = output_group_name(
source_group=sample.source_group,
actual_group=sample.actual_group,
is_synthetic=sample.is_synthetic,
mode=args.output_group_mode,
)
row = copiar_sample(
sample=sample,
dst_root=dst_root,
split_name=split_name,
out_group_name=out_group,
copy_meta_preview=not args.no_meta_preview,
copy_visuals=not args.no_visuals,
)
if row is not None:
row["decision_reason"] = reason
manifest_rows.append(
row
)
manifest_path = (
args.manifest
or os.path.join(
dst_root,
"split_manifest.csv",
)
)
audit_path = (
args.audit
or os.path.join(
dst_root,
"split_audit.csv",
)
)
families_path = (
args.families
or os.path.join(
dst_root,
"split_families.csv",
)
)
skipped_path = (
args.skipped
or os.path.join(
dst_root,
"split_skipped.csv",
)
)
summary_path = (
args.summary
or os.path.join(
dst_root,
"split_summary.json",
)
)
audit_rows = []
skipped_rows = []
for sample in analyzed:
split_name, reason = decisions[
id(sample)
]
row = sample_audit_row(
sample,
split_name,
reason,
)
audit_rows.append(
row
)
if split_name is None:
skipped_rows.append(
row
)
for sample in rejected:
audit_rows.append(
sample_audit_row(
sample,
None,
"audit_rejected",
)
)
if audit_rows:
write_csv(
audit_path,
audit_rows,
list(audit_rows[0].keys()),
)
family_rows = family_audit_rows(
units
)
if family_rows:
write_csv(
families_path,
family_rows,
list(family_rows[0].keys()),
)
if skipped_rows:
write_csv(
skipped_path,
skipped_rows,
list(skipped_rows[0].keys()),
)
if (
not args.dry_run
and manifest_rows
):
write_csv(
manifest_path,
manifest_rows,
list(manifest_rows[0].keys()),
)
summary = {
"schema": "oak_fcc3_agri_teacher_v2_split_v1",
"resolution": list(resolution),
"strategy": args.strategy,
"src_roots": src_roots,
"train_only_roots": train_only_roots,
"dst_root": dst_root,
"proportions": {
"train": p_train,
"val": p_val,
"test": p_test,
},
"options": {
"use_tensor_features": args.use_tensor_features,
"tensor_feature_stride": args.tensor_feature_stride,
"merge_near_duplicate_families": args.merge_near_duplicate_families,
"near_dup_hamming": args.near_dup_hamming,
"near_dup_mask_l1": args.near_dup_mask_l1,
"synthetic_train_only": args.synthetic_train_only,
"synthetic_respect_family_split": args.synthetic_respect_family_split,
"output_group_mode": args.output_group_mode,
},
"teacher_bins": {
"target_tiny_frac": args.target_tiny_frac,
"target_small_frac": args.target_small_frac,
"target_medium_frac": args.target_medium_frac,
"near_radius_frac": args.near_radius_frac,
"close_weed_cane_threshold": args.close_weed_cane_threshold,
},
"counts": {
"samples_seen": len(samples),
"samples_analyzed": len(analyzed),
"samples_rejected": len(rejected),
"real_samples": len(real_analyzed),
"folder_mask_mismatches": len(mismatches),
"base_families": len(base_units),
"leakage_units": len(units),
"near_duplicate_families_merged": merged_count,
"family_split": family_counts,
"samples_split": totals,
},
"strata": stratum_counts,
"paths": {
"manifest": (
manifest_path
if not args.dry_run
else None
),
"audit": audit_path,
"families": families_path,
"skipped": skipped_path,
},
"dry_run": bool(args.dry_run),
}
save_json(
summary_path,
summary,
)
print("==========================================")
print("RESULTADO")
print(
f"Samples reais : {len(real_analyzed)}"
)
print(
f"Mismatch pasta/GT : {len(mismatches)}"
)
print(
f"Famílias base : {len(base_units)}"
)
print(
f"Unidades leakage : {len(units)}"
)
print(
f"Near-dups mescladas: {merged_count}"
)
print(
f"TRAIN : {totals['train']} "
f"(real={totals['real_train']} synthetic={totals['synthetic_train']})"
)
print(
f"VAL real : {totals['real_val']}"
)
print(
f"TEST real : {totals['real_test']}"
)
print(
f"Skipped : {totals['skipped']}"
)
print(f"Audit : {audit_path}")
print(f"Families : {families_path}")
print(f"Summary : {summary_path}")
if not args.dry_run:
print(f"Manifest : {manifest_path}")
print("==========================================")
if args.dry_run:
print("✅ DRY-RUN concluído. Nenhum dataset foi copiado.")
else:
print("✅ Split agrícola Teacher V2 concluído.")
if __name__ == "__main__":
main()