# 👉 Treinador robusto para FastSCNN com ROI, pesos dinâmicos, Dice+CE, EMA, acumulação de gradiente e early-stop bicritério # Compatível com a sua estrutura atual (FastSCNN, ROISegDataset, utils.compute_roi_indices / carregar_labelmap_completo) # Principais melhorias vs _8_train_fastscnn.py: # - Corrige bug do argumento --use-weights (antes "use_wights") # - Usa de fato os pesos também no Dice (opcional, com normalização) # - Semeadura determinística opcional + cudnn.benchmark para performance # - Gradient clipping, acumulação de gradiente e AMP habilitáveis por CLI # - EMA (Exponential Moving Average) do modelo para validação/checkout # - Sampler que garante classes por batch (igual ao seu), mas com opções por CLI # - Scheduler Cosine → Plateau com aquecimento (warmup) do LR e dos pesos das classes # - Checkpoints por menor val_loss e por melhor F1 da classe principal # - Plots e logs semelhantes, mantendo compatibilidade de caminhos """ python _8_train_fastscnn_v2.py --epochs 120 --batch 16 --accum 2 --workers 4 --amp --lr 3e-4 --warmup-epochs 5 --min-lr 1e-5 --use-weights --cw-alpha 1.1 --cw-erva-min 1.0 --cw-cana-max 2.2 --cw-chao-min 0.6 --dice-max 0.15 --ensure-per-batch "erva,cana" --es-classes "erva,cana" --es-agg macro --patience-es 12 --delta-es-min 2e-4 --plateau-patience 4 --ema --tversky-classes "cana" --tversky-weights "0.10" --tversky-alphas "0.8" --tversky-betas "0.2" python _8_train_fastscnn_v2.py --epochs 8 --batch 16 --accum 2 --workers 4 --amp --lr 1.5e-4 --min-lr 1e-5 --warmup-epochs 1 --use-weights --cw-alpha 1.2 --cw-erva-min 1.8 --cw-cana-max 1.2 --cw-chao-min 0.6 --dice-max 0.12 --tversky-classes "erva,cana" --tversky-weights "0.40,0.08" --tversky-alphas "0.65,0.25" --tversky-betas "0.35,0.75" --ensure-per-batch "erva" --es-classes "erva,cana" --es-agg harmonic --patience-es 10 --delta-es-min 2e-4 --ema --checkpoint .\oak-1\backup\fast_scnn\ervas_es\ervas_es_best_f1_erva_checkpoint.pth python _8_train_fastscnn_v2.py --epochs 12 --batch 16 --accum 2 --workers 4 --amp --lr 9e-5 --min-lr 6e-6 --warmup-epochs 0 --use-weights --cw-alpha 1.1 --cw-erva-min 1.6 --cw-cana-max 1.5 --cw-chao-min 0.6 --dice-max 0.12 --tversky-classes "erva,cana" --tversky-weights "0.32,0.14" --tversky-alphas "0.70,0.35" --tversky-betas "0.30,0.65" --ensure-per-batch "erva,cana" --es-classes "erva,cana" --es-agg harmonic --patience-es 10 --delta-es-min 1e-4 --ema --checkpoint .\oak-1\backup\fast_scnn\ervas_es\ervas_es_best_es_harmonic_erva_cana_checkpoint.pth python _8_train_fastscnn_v2.py --epochs 12 --batch 16 --accum 2 --workers 4 --amp --lr 6e-5 --min-lr 5e-6 --warmup-epochs 0 --use-weights --cw-alpha 1.1 --cw-erva-min 1.5 --cw-cana-max 1.6 --cw-chao-min 0.6 --dice-max 0.12 --tversky-classes "erva,cana" --tversky-weights "0.30,0.14" --tversky-alphas "0.70,0.40" --tversky-betas "0.30,0.60" --ensure-per-batch "erva,cana" --es-classes "erva,cana" --es-agg macro --patience-es 10 --delta-es-min 1e-4 --ema --checkpoint .\oak-1\backup\fast_scnn\ervas_es\ervas_es_best_es_harmonic_erva_cana_checkpoint.pth python _8_train_fastscnn_v2.py --epochs 10 --batch 16 --accum 2 --workers 4 --amp --lr 5e-5 --min-lr 5e-6 --warmup-epochs 0 --use-weights --cw-alpha 1.15 --cw-erva-min 1.7 --cw-cana-max 1.5 --cw-chao-min 0.55 --dice-max 0.16 --tversky-classes "erva,cana" --tversky-weights "0.34,0.12" --tversky-alphas "0.35,0.60" --tversky-betas "0.65,0.40" --ensure-per-batch "erva,cana" --es-classes "erva,cana" --es-agg macro --patience-es 10 --delta-es-min 1e-4 --ema --checkpoint .\oak-1\backup\fast_scnn\ervas_es\ervas_es_best_es_macro_erva_cana_checkpoint.pth python _8_train_fastscnn_v2.py --epochs 10 --batch 16 --accum 2 --workers 4 --amp --lr 4e-5 --min-lr 4e-6 --warmup-epochs 0 --use-weights --cw-alpha 1.15 --cw-erva-min 1.9 --cw-cana-max 1.35 --cw-chao-min 0.55 --dice-max 0.20 --tversky-classes "erva,cana" --tversky-weights "0.36,0.10" --tversky-alphas "0.30,0.65" --tversky-betas "0.70,0.35" --ensure-per-batch "erva,cana" --es-classes "erva,cana" --es-agg macro --patience-es 10 --delta-es-min 1e-4 --ema --checkpoint .\oak-1\backup\fast_scnn\ervas_es\ervas_es_best_es_macro_erva_cana_checkpoint.pth python _8_train_fastscnn_v2.py --epochs 8 --batch 16 --accum 2 --workers 4 --amp --lr 3e-5 --min-lr 3e-6 --warmup-epochs 0 --use-weights --cw-alpha 1.15 --cw-erva-min 2.0 --cw-cana-max 1.30 --cw-chao-min 0.55 --dice-max 0.22 --tversky-classes "erva,cana" --tversky-weights "0.38,0.09" --tversky-alphas "0.25,0.70" --tversky-betas "0.75,0.30" --ensure-per-batch "erva,cana" --es-classes "erva,cana" --es-agg macro --patience-es 10 --delta-es-min 1e-4 --ema --checkpoint .\oak-1\backup\fast_scnn\ervas_es\ervas_es_best_es_macro_erva_cana_checkpoint.pth """ # A partir do novo dataset """ python _8_train_fastscnn_v2.py --epochs 120 --batch 16 --accum 2 --workers 4 --amp --lr 1e-4 --min-lr 5e-6 --warmup-epochs 1 --use-weights --cw-alpha 1.15 --cw-erva-min 1.9 --cw-cana-max 1.35 --cw-chao-min 0.55 --dice-max 0.20 --tversky-classes "erva,cana" --tversky-weights "0.36,0.10" --tversky-alphas "0.30,0.65" --tversky-betas "0.70,0.35" --ensure-per-batch "erva,cana" --es-classes "erva,cana" --es-agg macro --patience-es 12 --delta-es-min 2e-4 --plateau-patience 3 --ema --checkpoint .\oak-1\backup\fast_scnn\ervas_es\ervas_es_best_es_macro_erva_cana_checkpoint_bkp.pth python _8_train_fastscnn_v2.py --epochs 30 --batch 16 --accum 2 --workers 4 --amp --lr 7e-5 --min-lr 5e-6 --warmup-epochs 0 --use-weights --cw-alpha 1.12 --cw-erva-min 1.7 --cw-cana-max 1.6 --cw-chao-min 0.55 --dice-max 0.20 --tversky-classes "erva,cana" --tversky-weights "0.32,0.16" --tversky-alphas "0.35,0.45" --tversky-betas "0.65,0.55" --ensure-per-batch "erva,cana" --es-classes "erva,cana" --es-agg harmonic --patience-es 10 --delta-es-min 2e-4 --plateau-patience 2 --ema --checkpoint .\oak-1\backup\fast_scnn\ervas_es\ervas_es_best_es_macro_erva_cana_checkpoint_bkp.pth """ """ python _8_train_fastscnn_v2.py --epochs 50 --batch 12 --accum 2 --workers 4 --amp --lr 1.0e-4 --min-lr 8e-6 --ensure-per-batch "erva,cana" --use-weights --cw-alpha 1.15 --tversky-classes "erva,cana" --tversky-weights "0.40,0.60" --tversky-alphas "0.35,0.35" --tversky-betas "0.65,0.65" --dice-max 0.15 --es-classes "erva,cana" --es-agg harmonic --patience-es 10 --delta-es-min 2e-4 --ema python _8_train_fastscnn_v2.py --epochs 20 --batch 12 --accum 2 --workers 4 --amp --lr 1.2e-4 --min-lr 1e-5 --ensure-per-batch "cana" --use-weights --cw-alpha 1.35 --tversky-classes "cana" --tversky-weights "1.0" --tversky-alphas "0.20" --tversky-betas "0.80" --dice-max 0.10 --es-classes "cana" --es-agg max --patience-es 6 --delta-es-min 1e-4 --ema python _8_train_fastscnn_v2.py --epochs 40 --batch 12 --accum 2 --workers 4 --amp --lr 1.2e-4 --min-lr 1e-5 --ensure-per-batch "cana" --use-weights --cw-alpha 1.60 --tversky-classes "cana" --tversky-weights "1.0" --tversky-alphas "0.10" --tversky-betas "0.90" --dice-max 0.10 --es-classes "cana" --es-agg max --patience-es 8 --delta-es-min 1e-4 --ema python _8_train_fastscnn_v2.py --epochs 25 --batch 12 --accum 2 --workers 4 --amp --lr 8e-5 --min-lr 6e-6 --ensure-per-batch "erva,cana" --use-weights --cw-alpha 1.25 --tversky-classes "erva,cana" --tversky-weights "0.35,0.65" --tversky-alphas "0.25,0.20" --tversky-betas "0.75,0.80" --dice-max 0.12 --es-classes "erva,cana" --es-agg harmonic --patience-es 10 --delta-es-min 2e-4 --ema """ """ python _8_train_fastscnn_v2.py --epochs 80 --batch 4 --accum 4 --workers 4 --lr 1e-4 --min-lr 1e-6 --amp --ema --ema-decay 0.999 --use-weights --cw-alpha 1.08 --cw-bgfloor 0.6 --dice-max 0.12 --warmup-epochs 5 --plateau-patience 6 --plateau-factor 0.5 --plateau-cooldown 1 --es-classes "erva,cana" --es-agg harmonic --delta-es-min 1e-4 --patience-es 8 --ensure-per-batch "erva,cana" --cudnn-benchmark --seed 42 python _8_train_fastscnn_v2.py --epochs 20 --batch 4 --accum 4 --workers 4 --lr 6e-5 --min-lr 6e-6 --amp --ema --ema-decay 0.999 --use-weights --cw-alpha 1.12 --cw-cana-max 6.0 --cw-bgfloor 0.45 --dice-max 0.25 --tversky-classes "cana,erva" --tversky-weights "1.10,0.05" --tversky-alphas "0.20,0.55" --tversky-betas "0.80,0.45" --tversky-in-val --ensure-per-batch "cana" --es-classes "cana" --es-agg macro --delta-es-min 2e-4 --patience-es 6 --warmup-epochs 0 --plateau-patience 2 --plateau-factor 0.5 --plateau-cooldown 1 --checkpoint python _8_train_fastscnn_v2.py --epochs 12 --batch 4 --accum 4 --workers 4 --lr 4e-5 --min-lr 6e-6 --amp --ema --ema-decay 0.999 --use-weights --cw-alpha 1.14 --cw-cana-max 7.5 --cw-bgfloor 0.40 --dice-max 0.30 --tversky-classes "cana,erva" --tversky-weights "1.30,0.05" --tversky-alphas "0.15,0.55" --tversky-betas "0.85,0.45" --tversky-in-val --ensure-per-batch "cana" --es-classes "cana" --es-agg macro --delta-es-min 2e-4 --patience-es 6 --plateau-patience 2 --plateau-factor 0.5 --plateau-cooldown 1 --checkpoint python _8_train_fastscnn_v2.py --epochs 18 --batch 4 --accum 4 --workers 4 --lr 2.5e-5 --min-lr 6e-6 --amp --ema --ema-decay 0.999 --use-weights --cw-alpha 1.18 --cw-cana-max 9.0 --cw-bgfloor 0.35 --dice-max 0.35 --tversky-classes "cana,erva" --tversky-weights "1.60,0.05" --tversky-alphas "0.10,0.55" --tversky-betas "0.90,0.45" --tversky-in-val --ensure-per-batch "cana" --es-classes "cana" --es-agg macro --delta-es-min 2e-4 --patience-es 8 --plateau-patience 2 --plateau-factor 0.6 --plateau-cooldown 1 --checkpoint python _8_train_fastscnn_v2.py --epochs 12 --batch 4 --accum 4 --workers 4 --lr 1.8e-5 --min-lr 6e-6 --amp --ema --ema-decay 0.999 --use-weights --cw-alpha 1.20 --cw-cana-max 10.5 --cw-bgfloor 0.32 --dice-max 0.35 --tversky-classes "cana,erva" --tversky-weights "1.80,0.10" --tversky-alphas "0.08,0.55" --tversky-betas "0.92,0.45" --tversky-in-val --ensure-per-batch "cana,erva" --es-classes "cana" --es-agg macro --delta-es-min 2e-4 --patience-es 8 --plateau-patience 2 --plateau-factor 0.6 --plateau-cooldown 1 --checkpoint python _8_train_fastscnn_v2.py --epochs 10 --batch 4 --accum 4 --workers 4 --lr 1.5e-5 --min-lr 6e-6 --amp --ema --ema-decay 0.999 --use-weights --cw-alpha 1.06 --cw-cana-max 6.0 --cw-bgfloor 0.52 --dice-max 0.25 --tversky-classes "cana,erva" --tversky-weights "1.00,0.20" --tversky-alphas "0.35,0.55" --tversky-betas "0.65,0.45" --tversky-in-val --ensure-per-batch "cana,erva" --es-classes "erva,cana" --es-agg harmonic --delta-es-min 2e-4 --patience-es 8 --plateau-patience 2 --plateau-factor 0.6 --plateau-cooldown 1 --checkpoint python _8_train_fastscnn_v2.py --epochs 8 --batch 4 --accum 4 --workers 4 --lr 1.2e-5 --min-lr 6e-6 --amp --ema --ema-decay 0.999 --use-weights --cw-alpha 1.04 --cw-cana-max 5.5 --cw-bgfloor 0.50 --dice-max 0.22 --tversky-classes "cana,erva" --tversky-weights "1.00,0.25" --tversky-alphas "0.40,0.55" --tversky-betas "0.60,0.45" --tversky-in-val --ensure-per-batch "cana,erva" --es-classes "erva,cana" --es-agg harmonic --delta-es-min 2e-4 --patience-es 8 --plateau-patience 2 --plateau-factor 0.7 --plateau-cooldown 1 --checkpoint .\oak-1\backup\fast_scnn\ervas_res\ervas_res_best_es_harmonic_erva_cana_checkpoint.pth python _8_train_fastscnn_v2.py --epochs 8 --batch 4 --accum 4 --workers 4 --lr 8.4e-06 --min-lr 6e-06 --amp --ema --ema-decay 0.999 --use-weights --cw-alpha 1.04 --cw-cana-max 5.8 --cw-bgfloor 0.50 --dice-max 0.22 --tversky-classes "cana,erva" --tversky-weights "1.10,0.25" --tversky-alphas "0.38,0.55" --tversky-betas "0.62,0.45" --tversky-in-val --ensure-per-batch "cana,erva" --es-classes "erva,cana" --es-agg harmonic --delta-es-min 2e-4 --patience-es 6 --plateau-patience 2 --plateau-factor 0.7 --plateau-cooldown 1 --checkpoint .\oak-1\backup\fast_scnn\ervas_res\ervas_res_best_es_harmonic_erva_cana_checkpoint.pth python _8_train_fastscnn_v2.py --epochs 6 --batch 4 --accum 4 --workers 4 --lr 7.5e-06 --min-lr 6e-06 --amp --ema --ema-decay 0.999 --use-weights --cw-alpha 1.04 --cw-cana-max 5.5 --cw-bgfloor 0.53 --dice-max 0.22 --tversky-classes "cana,erva" --tversky-weights "1.00,0.25" --tversky-alphas "0.46,0.55" --tversky-betas "0.54,0.45" --tversky-in-val --ensure-per-batch "cana,erva" --es-classes "erva,cana" --es-agg harmonic --delta-es-min 2e-4 --patience-es 6 --plateau-patience 2 --plateau-factor 0.7 --plateau-cooldown 1 --checkpoint .\oak-1\backup\fast_scnn\ervas_res\ervas_res_best_es_harmonic_erva_cana_checkpoint.pth """ import itertools import math import os import random os.environ["MPLBACKEND"] = "Agg" # segurança extra import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt plt.ioff() import json import time import argparse from PIL import Image import numpy as np import torch import torch.nn.functional as F import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from collections import OrderedDict from fast_scnn import FastSCNN from roi_seg_dataset import ROISegDataset from utils import carregar_labelmap_completo, compute_roi_indices # ============================= # Utilidades # ============================= def set_seed(seed: int = 42): if seed is None: return random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = False def maybe_benchmark(use_benchmark: bool = True): # Em GPUs modernas, benchmark True acelera convs com shapes fixos torch.backends.cudnn.benchmark = bool(use_benchmark) def _infer_ignore_id(ignore_rgb, default_id=255): import numpy as _np if isinstance(ignore_rgb, (list, tuple)): if len(ignore_rgb) == 1 and isinstance(ignore_rgb[0], (int, _np.integer)): return int(ignore_rgb[0]) if len(ignore_rgb) == 3: return default_id if isinstance(ignore_rgb, (int, _np.integer)): return int(ignore_rgb) return default_id @torch.no_grad() def confmat_update(confmat, pred, target, num_classes, ignore_index=None): # pred, target: (B,H,W) if ignore_index is not None: mask = target != ignore_index target = target[mask] pred = pred[mask] k = (target * num_classes + pred).to(torch.int64) binc = torch.bincount(k, minlength=num_classes**2) confmat += binc.reshape(num_classes, num_classes) return confmat def metrics_from_confmat(confmat, main_class_id=None): # confmat: CxC (linhas = GT, colunas = Pred) cm = confmat.float() eps = 1e-7 tp = torch.diag(cm) fp = cm.sum(0) - tp fn = cm.sum(1) - tp # Por-classe precision_per_class = tp / (tp + fp + eps) recall_per_class = tp / (tp + fn + eps) f1_per_class = 2 * precision_per_class * recall_per_class / (precision_per_class + recall_per_class + eps) # IoU por classe denom_iou = tp + fp + fn + eps iou_per_class = tp / denom_iou # mIoU: média só das classes presentes no GT deste batch (tp+fn > 0) present_mask = (tp + fn) > 0 miou = iou_per_class[present_mask].mean().item() if present_mask.any() else 0.0 # Pixel accuracy (micro) pix_acc = tp.sum() / (cm.sum() + eps) # Métricas da classe principal (se pedida) main_class_metrics = None if main_class_id is not None and 0 <= main_class_id < cm.shape[0]: main_class_metrics = { "precision": precision_per_class[main_class_id].item(), "recall": recall_per_class[main_class_id].item(), "f1": f1_per_class[main_class_id].item(), "iou": iou_per_class[main_class_id].item(), } return { "miou": miou, "pixel_acc": pix_acc.item(), "iou_per_class": iou_per_class.cpu().tolist(), "precision_per_class": precision_per_class.cpu().tolist(), "recall_per_class": recall_per_class.cpu().tolist(), "f1_per_class": f1_per_class.cpu().tolist(), "tp_per_class": tp.cpu().tolist(), "fp_per_class": fp.cpu().tolist(), "fn_per_class": fn.cpu().tolist(), "present_mask": present_mask.cpu().tolist(), # opcional, ajuda a depurar "main_class": main_class_metrics, } def _id_by_name(d, name): name = str(name).lower() if isinstance(d, dict): for cid, nm in d.items(): if isinstance(nm, str) and name in nm.lower(): return cid elif isinstance(d, (list, tuple)): for i, nm in enumerate(d): if isinstance(nm, str) and name in nm.lower(): return i return None def _ids_by_names(CLASS_NAMES, names_csv): names = _parse_csv_strs(names_csv) norm = _normalize_class_names(CLASS_NAMES, len(CLASS_NAMES)) ids = [] for nm in names: cid = _id_by_name(norm, nm) if cid is not None: ids.append(cid) return ids def _aggregate(vals, mode="macro"): vals = [float(v) for v in vals if v is not None] if not vals: return None if mode == "macro": return sum(vals) / len(vals) if mode == "min": return min(vals) if mode == "harmonic": eps = 1e-8 return len(vals) / sum(1.0 / (v + eps) for v in vals) return sum(vals) / len(vals) def _to_list(x): if x is None: return None if torch.is_tensor(x): return x.detach().cpu().tolist() return list(x) def _normalize_class_names(class_names, n_classes): """ Aceita: - list/tuple -> usa direto (ajusta tamanho se preciso) - dict id->name (chaves int) -> ordena por id - dict name->id (chaves str) -> inverte p/ lista por id - None -> ['0','1',...] """ if class_names is None: return [str(i) for i in range(n_classes)] # já é lista/tupla if isinstance(class_names, (list, tuple)): names = list(class_names) # ajusta tamanho se vier curto/long if len(names) < n_classes: names += [f"{len(names)+i}" for i in range(n_classes - len(names))] return names[:n_classes] # é dict if isinstance(class_names, dict): # id->name if all(isinstance(k, int) for k in class_names.keys()): return [class_names.get(i, str(i)) for i in range(n_classes)] # name->id names = [f"{i}" for i in range(n_classes)] for name, idx in class_names.items(): if isinstance(idx, int) and 0 <= idx < n_classes: names[idx] = str(name) return names # fallback return [str(i) for i in range(n_classes)] def _format_all_classes(m, class_names, main_idx): # pega vetores f1s = _to_list(m.get('f1_per_class')) ious = _to_list(m.get('iou_per_class')) precs = _to_list(m.get('precision_per_class')) recs = _to_list(m.get('recall_per_class')) tps = _to_list(m.get('tp_per_class')) fps = _to_list(m.get('fp_per_class')) fns = _to_list(m.get('fn_per_class')) # descobre n_classes lens = [len(x) for x in [f1s, ious, precs, recs, tps, fps, fns] if x is not None] n = max(lens) if lens else 0 names = _normalize_class_names(class_names, n) # totais p/ shares (usa total de pixels de GT como base) total_gt = 0.0 if tps is not None and fns is not None: total_gt = float(sum((tp + fn) for tp, fn in zip(tps, fns))) total_gt = total_gt if total_gt > 0 else 1.0 parts = [] for i in range(n): name = names[i] tag = "🌿" if (main_idx is not None and i == main_idx) else "" f1 = f1s[i] if (f1s is not None and i < len(f1s)) else None iou = ious[i] if (ious is not None and i < len(ious)) else None pr = precs[i] if (precs is not None and i < len(precs)) else None rc = recs[i] if (recs is not None and i < len(recs)) else None tp = tps[i] if (tps is not None and i < len(tps)) else None fp = fps[i] if (fps is not None and i < len(fps)) else None fn = fns[i] if (fns is not None and i < len(fns)) else None # shares (% de pixels GT dessa classe e % preditos nessa classe) gt_share = ((tp + fn) / total_gt) if (tp is not None and fn is not None) else None pred_share = ((tp + fp) / total_gt) if (tp is not None and fp is not None) else None def _fmt(x, nd=4): try: if x is None or (isinstance(x, float) and (math.isnan(x) or math.isinf(x))): return "—" return f"{float(x):.{nd}f}" except Exception: return "—" s = (f"{tag}{name}: " f"F1={_fmt(f1)} IoU={_fmt(iou)} " f"P={_fmt(pr)} R={_fmt(rc)} " f"GT%={_fmt(100.0*gt_share, nd=1)} Pred%={_fmt(100.0*pred_share, nd=1)}") parts.append(s) return " | " + " ".join(parts) def _parse_csv_floats(s): if not s: return [] out = [] for tok in s.split(","): t = tok.strip() if t: out.append(float(t)) return out def _parse_csv_strs(s): if not s: return [] return [tok.strip() for tok in s.split(",") if tok.strip()] def tversky_loss_one(logits, target, class_id, alpha=0.5, beta=0.5, ignore_index=255, eps=1e-6): """ Tversky para UMA classe: 1 - TP / (TP + α*FP + β*FN) logits: [N,C,H,W], target: [N,H,W] """ N, C, H, W = logits.shape # prob da classe p = torch.softmax(logits, dim=1)[:, class_id, ...] # [N,H,W] valid = (target != ignore_index) tgt = (target == class_id).float() p = p * valid tgt = tgt * valid tp = (p * tgt).sum() fp = (p * (1.0 - tgt)).sum() fn = ((1.0 - p) * tgt).sum() ti = tp / (tp + alpha * fp + beta * fn + eps) return 1.0 - ti def build_tversky_cfg(args, class_names): """ Lê args e cria uma lista de dicts: [{"name":str, "id":int, "weight":float, "alpha":float, "beta":float}, ...] Suporta broadcast quando weights/alphas/betas têm 1 elemento. """ cls_names = _parse_csv_strs(args.tversky_classes) if not cls_names: return [] ws = _parse_csv_floats(args.tversky_weights) or [0.2] # default weight al = _parse_csv_floats(args.tversky_alphas) or [0.5] be = _parse_csv_floats(args.tversky_betas) or [0.5] def bcast(v, n): return v * n if len(v) == 1 else v ws = bcast(ws, len(cls_names)) al = bcast(al, len(cls_names)) be = bcast(be, len(cls_names)) if not (len(ws) == len(al) == len(be) == len(cls_names)): raise ValueError(f"Comprimentos inconsistentes: classes={len(cls_names)} " f"weights={len(ws)} alphas={len(al)} betas={len(be)}") norm_names = _normalize_class_names(class_names, len(class_names)) cfg = [] for name, w, a, b in zip(cls_names, ws, al, be): cid = _id_by_name(norm_names, name) if cid is None: print(f"⚠️ Tversky: classe '{name}' não encontrada nas classes do dataset — ignorando.") continue if w <= 0: # peso zero -> silenciosamente ignora continue cfg.append({"name": name, "id": int(cid), "weight": float(w), "alpha": float(a), "beta": float(b)}) return cfg def _merge_logits_eq_rest(logits, eq_ids, rest_ids): # logits: [N,C,H,W] l_eq = torch.logsumexp(logits[:, eq_ids, ...], dim=1, keepdim=True) l_rest = torch.logsumexp(logits[:, rest_ids, ...], dim=1, keepdim=True) return torch.cat([l_eq, l_rest], dim=1) # [N,2,H,W] def _make_y_bin(y, eq_ids, ignore_id): # 0 = EQ, 1 = REST; preserva ignore_id y_bin = torch.full_like(y, fill_value=1) for cid in eq_ids: y_bin = torch.where(y == cid, torch.zeros_like(y_bin), y_bin) if ignore_id is not None: y_bin = torch.where(y == ignore_id, torch.tensor(ignore_id, device=y.device, dtype=y.dtype), y_bin) return y_bin def _dice_from_probs(p, t, mask=None, eps=1e-6): # p,t: [N,H,W] em {0..1}; mask opcional booleana if mask is not None: p = p[mask]; t = t[mask] inter = (p * t).sum() denom = p.sum() + t.sum() return 1.0 - (2*inter + eps) / (denom + eps) # dice loss def _get_mask_roi_from_ds(ds, i, roi_inicio, roi_tamanho): mask_path = None if hasattr(ds, "mask_paths"): mask_path = ds.mask_paths[i] elif hasattr(ds, "items"): item = ds.items[i] if isinstance(item, dict) and "mask" in item: mask_path = item["mask"] if mask_path is not None: m = np.array(Image.open(mask_path).convert("L")) else: _, y = ds[i] m = y.cpu().numpy() H = m.shape[0] y_fim, y_ini = compute_roi_indices(H, roi_inicio, roi_tamanho) return m[y_fim:y_ini, :] def compute_presence_indices(ds, class_ids, roi_inicio, roi_tamanho): presence = {cid: [] for cid in class_ids} others = [] for i in range(len(ds)): roi = _get_mask_roi_from_ds(ds, i, roi_inicio, roi_tamanho) found_any = False for cid in class_ids: if (roi == cid).any(): presence[cid].append(i) found_any = True if not found_any: others.append(i) return presence, others class EnsureClassesBatchSampler(torch.utils.data.Sampler): """Garante >=1 amostra de cada classe requerida por batch.""" def __init__(self, presence, total_indices, batch_size, required_classes, seed=42): self.presence = presence self.required = [c for c in required_classes if len(presence.get(c, [])) > 0] self.batch_size = batch_size self.iters = {c: itertools.cycle(self.presence[c]) for c in self.required} self.rest_iter = itertools.cycle(list(total_indices)) self.rng = random.Random(seed) self._length = max(1, int(np.ceil(len(total_indices) / float(batch_size)))) def __iter__(self): for _ in range(self._length): batch = [] for c in self.required: batch.append(next(self.iters[c])) while len(batch) < self.batch_size: batch.append(next(self.rest_iter)) self.rng.shuffle(batch) yield batch def __len__(self): return self._length def compute_class_weights_from_split( train_split_root: str, labelmap_path: str, roi_inicio: float, roi_tamanho: float, *, alpha: float = 1.2, w_min: float = 0.3, w_max: float = 4.0, floor_bg: float = 0.4, max_samples_per_group: int = 300, ): """ Calcula pesos dinâmicos (median freq ^ alpha) SOBRE A ROI das máscaras do split/train. """ _, _, classes, ignore_rgb = carregar_labelmap_completo(labelmap_path) ignore_id = _infer_ignore_id(ignore_rgb, default_id=255) class_ids = sorted(classes.keys()) counts = np.zeros(len(class_ids), dtype=np.int64) group_root = os.path.join(train_split_root, "group") group_dirs = [] if os.path.isdir(group_root): for g in sorted(os.listdir(group_root)): mdir = os.path.join(group_root, g, "masks") if os.path.isdir(mdir): group_dirs.append(mdir) else: mdir = os.path.join(train_split_root, "masks") if os.path.isdir(mdir): group_dirs.append(mdir) for mdir in group_dirs: n = 0 for p in os.listdir(mdir): if not p.lower().endswith(".png"): continue m = np.array(Image.open(os.path.join(mdir, p)).convert("L")) H = m.shape[0] y_fim, y_ini = compute_roi_indices(H, roi_inicio, roi_tamanho) roi = m[y_fim:y_ini, :] for i, cid in enumerate(class_ids): if cid == ignore_id: continue counts[i] += int((roi == cid).sum()) n += 1 if n >= max_samples_per_group: break total = int(counts.sum()) if total == 0: print("⚠️ compute_class_weights_from_split: não encontrei pixels válidos; usando pesos [1,1,...].") return None, None freqs = counts / total nonzero = freqs[freqs > 0] base = np.median(nonzero) if nonzero.size > 0 else 1.0 weights_arr = np.zeros_like(freqs, dtype=np.float32) for i, f in enumerate(freqs): if f <= 0: w = w_max else: w = (base / f) ** alpha w = float(np.clip(w, w_min, w_max)) weights_arr[i] = w for i, cid in enumerate(class_ids): name = str(classes[cid]).lower() if ("chao" in name) or ("background" in name): weights_arr[i] = max(weights_arr[i], floor_bg) max_cid = max(class_ids) weights_full = np.ones(max_cid + 1, dtype=np.float32) for i, cid in enumerate(class_ids): weights_full[cid] = weights_arr[i] pretty = {int(cid): (str(classes[cid]), float(weights_full[cid]), float(freqs[i])) for i, cid in enumerate(class_ids)} return weights_full, pretty class ModelEMA: """EMA dos pesos do modelo (independente da arquitetura).""" def __init__(self, model: nn.Module, decay: float = 0.999): import copy self.decay = decay device = next(model.parameters()).device self.ema = copy.deepcopy(model).to(device) self.ema.eval() for p in self.ema.parameters(): p.requires_grad_(False) @torch.no_grad() def update(self, model: nn.Module): d = self.decay msd = model.state_dict() for k, v_ema in self.ema.state_dict().items(): v_model = msd[k].detach() v_ema.copy_(v_ema * d + (1. - d) * v_model) def dice_loss(logits, target, ignore_index=255, class_weights=None, eps=1e-6): N, C, H, W = logits.shape pred = F.softmax(logits, dim=1) valid = (target != ignore_index) target_clamped = torch.clamp(target, 0, C-1) one_hot = torch.zeros((N, C, H, W), device=logits.device, dtype=pred.dtype) one_hot.scatter_(1, target_clamped.unsqueeze(1), 1.0) valid = valid.unsqueeze(1) pred = pred * valid one_hot = one_hot * valid inter = (pred * one_hot).sum(dim=(0, 2, 3)) pred_sum = pred.sum(dim=(0, 2, 3)) tgt_sum = one_hot.sum(dim=(0, 2, 3)) dice = (2 * inter + eps) / (pred_sum + tgt_sum + eps) if class_weights is not None: w = torch.ones(C, device=logits.device, dtype=pred.dtype) w[:class_weights.numel()] = class_weights loss = 1.0 - (w * dice).sum() / (w.sum() + eps) else: loss = 1.0 - dice.mean() return loss def _safe_torch_load(path, map_location, try_weights_only=False): try: return torch.load(path, map_location=map_location, weights_only=try_weights_only) except TypeError: # versões de torch que não têm weights_only return torch.load(path, map_location=map_location) # ============================= # Treino principal # ============================= def train(args): set_seed(args.seed) maybe_benchmark(args.cudnn_benchmark) # Config base do seu ambiente with open("config.json", "r") as f: config = json.load(f) MODELO = config["camera"] MODEL_NAME = config["model_name"] RESOLUCAO = config["resolucao"] ROI_INICIO = config["roi_inicio"] ROI_TAMANHO = config["roi_tamanho"] MAIN_CLASS_NAME = str(config.get("main_class_name", "erva")).lower() save_path = os.path.join(MODELO, "backup", config["modelo"], MODEL_NAME) dataset_path = os.path.join(MODELO, "dataset") labelmap_path = os.path.join(dataset_path, "labelmap.txt") os.makedirs(save_path, exist_ok=True) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Device: {device}") # ================= Dataset ================= ds_train = ROISegDataset( os.path.join(dataset_path, "split", "train"), save_path, ROI_INICIO, ROI_TAMANHO, RESOLUCAO[0], RESOLUCAO[1], labelmap_path ) ds_val = ROISegDataset( os.path.join(dataset_path, "split", "val"), save_path, ROI_INICIO, ROI_TAMANHO, RESOLUCAO[0], RESOLUCAO[1], labelmap_path ) CLASS_NAMES = getattr(ds_train, "classes", None) # pode ser list OU dict ES_CLASS_IDS = _ids_by_names(CLASS_NAMES, args.es_classes) # ex.: ids de ['erva','cana'] if not ES_CLASS_IDS: print("⚠️ ES: nenhuma classe válida encontrada; ES será ignorado.") best_es = -1.0 no_imp_es = 0 TV_CFG = build_tversky_cfg(args, CLASS_NAMES) if TV_CFG: print({"tversky_cfg": TV_CFG}) else: print("Tversky OFF (sem classes ou pesos <= 0)") eq_ids = None rest_ids = None if args.eq_classes: eq_names = [s.strip() for s in args.eq_classes.split(',') if s.strip()] norm_names = _normalize_class_names(CLASS_NAMES, len(CLASS_NAMES)) eq_ids = sorted([_id_by_name(norm_names, nm) for nm in eq_names if _id_by_name(norm_names, nm) is not None]) # REST = todas as classes válidas que NÃO estão no grupo eq (exc. ignore_id) rest_ids = [i for i in range(len(CLASS_NAMES)) if (i not in eq_ids)] # id->nome (preferir do dataset; fallback labelmap) if hasattr(ds_train, "classes") and isinstance(ds_train.classes, dict) and len(ds_train.classes) > 0: id_to_name = {int(k): str(v) for k, v in ds_train.classes.items()} else: _, _, id_to_name, _ = carregar_labelmap_completo(labelmap_path) id_to_name = {int(k): str(v) for k, v in id_to_name.items()} name_to_id = {v.strip().lower(): k for k, v in id_to_name.items()} # Classe principal main_class_id = _id_by_name(id_to_name, MAIN_CLASS_NAME) if main_class_id is not None: print(f"🌿 Classe PRIMARIA: id={main_class_id}, nome='{id_to_name[main_class_id]}'") else: print("⚠️ Classe PRIMARIA não encontrada; F1 específico será ignorado.") # ----- Sampler ensure-per-batch (opcional) ----- req_ids = [] if args.ensure_per_batch: for nm in [s.strip().lower() for s in args.ensure_per_batch.split(",") if s.strip()]: cid = name_to_id.get(nm) if cid is None: matches = [k for k, v in id_to_name.items() if nm in v.lower()] if len(matches) == 1: cid = matches[0] elif len(matches) > 1: print(f"⚠️ '--ensure-per-batch {nm}': ambíguo entre {[id_to_name[m] for m in matches]}; ignorando este nome.") cid = None else: print(f"⚠️ '--ensure-per-batch {nm}': classe não encontrada nas classes {list(name_to_id.keys())}.") if cid is not None and cid not in req_ids: req_ids.append(cid) if len(req_ids) > 0: presence, _ = compute_presence_indices(ds_train, class_ids=req_ids, roi_inicio=ROI_INICIO, roi_tamanho=ROI_TAMANHO) total_indices = range(len(ds_train)) batch_sampler = EnsureClassesBatchSampler( presence=presence, total_indices=total_indices, batch_size=args.batch, required_classes=req_ids, seed=args.seed or 42, ) dl_train = DataLoader(ds_train, batch_sampler=batch_sampler, num_workers=args.workers, pin_memory=True, persistent_workers=bool(args.workers > 0)) else: dl_train = DataLoader(ds_train, batch_size=args.batch, shuffle=True, num_workers=args.workers, pin_memory=True, persistent_workers=bool(args.workers > 0)) dl_val = DataLoader(ds_val, batch_size=args.batch, shuffle=False, num_workers=args.workers, pin_memory=True, persistent_workers=bool(args.workers > 0)) num_classes = len(ds_train.classes) ignore_id = getattr(ds_train, "ignore_id", 255) # ================= Modelo / Opt / LR ================= model = FastSCNN(num_classes=num_classes).to(device) if args.freeze_bn: for m in model.modules(): if isinstance(m, (nn.BatchNorm2d, nn.SyncBatchNorm)): m.eval() for p in m.parameters(): p.requires_grad_(False) print("🧊 BatchNorm congelado.") optimizer = optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay) # Warmup do LR (linear nos primeiros warmup_epochs) def lr_warmup_factor(epoch): if args.warmup_epochs <= 0: return 1.0 return min(1.0, max(0.0, epoch / float(args.warmup_epochs))) from torch import amp scaler = amp.GradScaler('cuda', enabled=args.amp) # ================= Pesos por classe (ROI) ================= train_split_root = os.path.join(dataset_path, "split", "train") weights_np, debug_info = compute_class_weights_from_split( train_split_root=train_split_root, labelmap_path=labelmap_path, roi_inicio=ROI_INICIO, roi_tamanho=ROI_TAMANHO, alpha=args.cw_alpha, w_min=args.cw_min, w_max=args.cw_max, floor_bg=args.cw_bgfloor, max_samples_per_group=args.cw_max_per_group, ) if weights_np is None or not args.use_weights: weights_t = None print("⚠️ Pesos dinâmicos indisponíveis ou desabilitados; usando CE sem pesos.") else: import pprint pprint.pprint({"class_weights": debug_info}) weights_t = torch.tensor(weights_np, device=device) # Warmup dos pesos de classe (λ 0→1) + clamps amigos ones = torch.ones_like(weights_t) if weights_t is not None else None def make_epoch_weights(epoch): if weights_t is None: return None, None cw_lambda = min(1.0, max(0.0, (epoch - 1) / max(1, args.cw_warmup))) w_epoch = (1.0 - cw_lambda) * ones + cw_lambda * weights_t # clamps opcionais (antes da normalização do Dice) erva_id = _id_by_name(ds_train.classes, "erva") cana_id = _id_by_name(ds_train.classes, "cana") chao_id = _id_by_name(ds_train.classes, "chao") if erva_id is not None: w_epoch[erva_id] = torch.clamp(w_epoch[erva_id], min=args.cw_erva_min) if cana_id is not None: w_epoch[cana_id] = torch.clamp(w_epoch[cana_id], max=args.cw_cana_max) if chao_id is not None: w_epoch[chao_id] = torch.clamp(w_epoch[chao_id], min=args.cw_chao_min) dice_w = (w_epoch / w_epoch.mean()).detach() return w_epoch, dice_w # EMA opcional ema = ModelEMA(model, decay=args.ema_decay) if args.ema else None # ================= Loop ================= start_epoch = 1 num_epochs = args.epochs best_val_loss = float("inf") best_main_class_f1 = -1.0 train_loss_history, val_loss_history, lr_history = [], [], [] f1_history, miou_history = [], [] no_imp_loss = no_imp_f1 = 0 epochs_since_switch = 0 # Recarregar checkpoint opcional if args.checkpoint and os.path.exists(args.checkpoint): print(f"🔁 Carregando: {args.checkpoint}") ckpt = _safe_torch_load(args.checkpoint, map_location=device, try_weights_only=False) def _is_plain_statedict(obj): if isinstance(obj, (dict, OrderedDict)): # heurística: parece um state_dict de modelo (tensores nas folhas) return any(k.endswith(".weight") or k.endswith(".bias") for k in obj.keys()) return False start_epoch = 1 num_epochs = args.epochs if _is_plain_statedict(ckpt): model.load_state_dict(ckpt, strict=True) # recomeça “limpo”: novo optimizer com LR do args for g in optimizer.param_groups: g['lr'] = args.lr print("🔁 Iniciei do state_dict (ex: *_best_f1_erva.pth).") elif isinstance(ckpt, dict) and "model" in ckpt: model.load_state_dict(ckpt["model"], strict=True) # se quiser retomar “igual”, descomente as duas linhas abaixo: # optimizer.load_state_dict(ckpt.get("optimizer", optimizer.state_dict())) # scaler.load_state_dict(ckpt.get("scaler", scaler.state_dict())) # mas como é fine-tune, geralmente é melhor resetar o otimizador: for g in optimizer.param_groups: g['lr'] = args.lr start_epoch = ckpt.get("epoch", 0) + 1 # não é obrigatório usar num_epochs = start_epoch + args.epochs # se você usa EMA e quer começar a partir do EMA salvo: if args.ema and "ema" in ckpt: ema.ema.load_state_dict(ckpt["ema"]) if "best_val_loss" in ckpt: best_val_loss = ckpt["best_val_loss"] if "best_main_class_f1" in ckpt: best_main_class_f1 = ckpt["best_main_class_f1"] if "best_es" in ckpt: best_es = ckpt["best_es"] print("🔁 Iniciei do checkpoint completo (ex: *_best_checkpoint.pth).") else: print("⚠️ Formato de checkpoint não reconhecido; seguindo sem carregar pesos.") cosine = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs, eta_min=args.min_lr) plateau = torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, mode="min", factor=args.plateau_factor, patience=args.plateau_patience, cooldown=args.plateau_cooldown, min_lr=args.min_lr, verbose=True ) active_sched = "cosine" # mistura Dice/CE crescente 0→dice_max nas primeiras 10 épocas dice_max = args.dice_max for epoch in range(start_epoch, num_epochs + 1): t0 = time.time() model.train() # Ajuste de LR warmup (manual multiplicativo) if args.warmup_epochs > 0 and epoch <= args.warmup_epochs: wf = epoch / float(args.warmup_epochs) for g in optimizer.param_groups: g['lr'] = args.lr * wf # Pesos deste epoch w_epoch, dice_w = make_epoch_weights(epoch) ce = nn.CrossEntropyLoss(ignore_index=ignore_id, weight=w_epoch) # mistura de perdas dice_mix = min(dice_max, (epoch-1)/10.0 * dice_max) ce_mix = 1.0 - dice_mix running_train_loss = 0.0 optimizer.zero_grad(set_to_none=True) for step, (x, y) in enumerate(dl_train, start=1): x, y = x.to(device), y.to(device) with torch.amp.autocast('cuda', enabled=args.amp): logits = model(x) if eq_ids is not None: # ===== FASE 1: perda binária EQ vs REST ===== logits2 = _merge_logits_eq_rest(logits, eq_ids, rest_ids) # [N,2,H,W] y_bin = _make_y_bin(y, eq_ids, ignore_id) # CE binária sobre logits2 ce_val = F.cross_entropy(logits2, y_bin, ignore_index=ignore_id) # Dice (ou Tversky) sobre probabilidade somada da EQ if dice_mix > 0: p_eq = logits2.softmax(1)[:, 0, ...] # prob da EQ t_eq = (y_bin == 0).float() # máscara válida (ignora pixels ignore_id) mask_valid = (y != ignore_id) if ignore_id is not None else None dloss = _dice_from_probs(p_eq, t_eq, mask_valid) d_val = dloss else: d_val = 0.0 tv_sum = 0.0 # (opcional: implementar tversky_from_probs aqui) else: ce_val = ce(logits, y) dloss = dice_loss(logits, y, ignore_index=ignore_id, class_weights=dice_w) d_val = dloss if dice_mix > 0 else 0.0 tv_sum = 0.0 for itv in TV_CFG: tv_sum = tv_sum + itv["weight"] * tversky_loss_one( logits, y, itv["id"], alpha=itv["alpha"], beta=itv["beta"], ignore_index=ignore_id ) loss = ce_mix * ce_val + dice_mix * d_val + tv_sum scaler.scale(loss).backward() if step % args.accum == 0: if args.grad_clip > 0: scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip) scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_none=True) if ema: ema.update(model) running_train_loss += loss.item() * x.size(0) * max(1, args.accum) avg_train_loss = running_train_loss / len(ds_train) train_loss_history.append(avg_train_loss) # ================= Validação ================= def evaluate(mdl: nn.Module, eq_ids): mdl.eval() running_val_loss = 0.0 confmat = torch.zeros((num_classes, num_classes), dtype=torch.int64, device=device) # extra: confusão 2x2 VEG vs REST (somente se eq_ids) conf_veg = torch.zeros((2,2), dtype=torch.int64, device=device) if eq_ids is not None else None with torch.no_grad(): for x, y in dl_val: x, y = x.to(device), y.to(device) with torch.amp.autocast('cuda', enabled=args.amp): logits = mdl(x) if eq_ids is not None: logits2 = _merge_logits_eq_rest(logits, eq_ids, rest_ids) y_bin = _make_y_bin(y, eq_ids, ignore_id) ce_val = F.cross_entropy(logits2, y_bin, ignore_index=ignore_id) if dice_mix > 0: p_eq = logits2.softmax(1)[:, 0, ...] t_eq = (y_bin == 0).float() mask_valid = (y != ignore_id) if ignore_id is not None else None d_val = _dice_from_probs(p_eq, t_eq, mask_valid) else: d_val = 0.0 tv_sum = 0.0 loss = ce_mix * ce_val + dice_mix * d_val + tv_sum else: ce_val = ce(logits, y) dloss = dice_loss(logits, y, ignore_index=ignore_id, class_weights=dice_w) d_val = dloss if dice_mix > 0 else 0.0 tv_sum = 0.0 if args.tversky_in_val and TV_CFG: for itv in TV_CFG: tv_sum = tv_sum + itv["weight"] * tversky_loss_one( logits, y, itv["id"], alpha=itv["alpha"], beta=itv["beta"], ignore_index=ignore_id ) loss = ce_mix * ce_val + dice_mix * d_val + tv_sum running_val_loss += loss.item() * x.size(0) pred = logits.argmax(1) confmat = confmat_update(confmat, pred, y, num_classes, ignore_index=ignore_id) # ==== Métrica da UNIÃO (veg = erva ∨ cana) ==== if args.eq_classes: # ids das classes do grupo (ex.: erva,cana) eq_names = [s.strip() for s in args.eq_classes.split(",") if s.strip()] cls_norm = _normalize_class_names(CLASS_NAMES, len(CLASS_NAMES)) eq_ids = [_id_by_name(cls_norm, n) for n in eq_names] valid = (y != ignore_id) # GT: 1 se for qualquer classe do grupo, 0 caso contrário y_veg = torch.zeros_like(y, dtype=torch.bool) for cid in eq_ids: y_veg |= (y == cid) y_veg &= valid # Pred: idem, 1 se previu qualquer classe do grupo p_veg = torch.zeros_like(pred, dtype=torch.bool) for cid in eq_ids: p_veg |= (pred == cid) p_veg &= valid # compacta só os válidos para evitar contar ignore como negativo yb = y_veg[valid].long() pb = p_veg[valid].long() idx = yb * 2 + pb # [0..3] = (TN, FP, FN, TP) na ordem abaixo b = torch.bincount(idx.view(-1), minlength=4) # acumula em um confusion 2x2 da união ao longo dos batches # crie antes do loop: conf_veg = torch.zeros((2,2), dtype=torch.long, device=device) conf_veg[0,0] += b[0] # TN conf_veg[0,1] += b[1] # FP conf_veg[1,0] += b[2] # FN conf_veg[1,1] += b[3] # TP avg_val = running_val_loss / len(ds_val) mets = metrics_from_confmat(confmat, main_class_id) if args.eq_classes: TP = conf_veg[1,1].float(); FP = conf_veg[0,1].float() FN = conf_veg[1,0].float() prec = TP / (TP + FP + 1e-9) rec = TP / (TP + FN + 1e-9) f1 = 2 * prec * rec / (prec + rec + 1e-9) iou = TP / (TP + FP + FN + 1e-9) mets["veg_union"] = {"precision": float(prec), "recall": float(rec), "f1": float(f1), "iou": float(iou)} return avg_val, mets # avalia com EMA se disponível if ema: # usa os pesos do EMA para validação avg_val_loss, m = evaluate(ema.ema, eq_ids) else: avg_val_loss, m = evaluate(model, eq_ids) # ===== Métricas pós-eval ===== val_loss_history.append(avg_val_loss) miou_history.append(m["miou"]) cur_lr = optimizer.param_groups[0]['lr'] lr_history.append(cur_lr) # F1 principal (se ainda quiser logar) main_class_f1 = m["main_class"]["f1"] if (m.get("main_class") is not None) else None if main_class_f1 is not None: f1_history.append(main_class_f1) # ES metric (agregado nas classes escolhidas) es_metric = None if args.eq_classes and m.get("veg_union"): es_metric = float(m["veg_union"]["f1"]) elif ES_CLASS_IDS and "f1_per_class" in m and m["f1_per_class"] is not None: es_vals = [float(m["f1_per_class"][cid]) for cid in ES_CLASS_IDS if cid < len(m["f1_per_class"])] es_metric = _aggregate(es_vals, mode=args.es_agg) elapsed = time.time() - t0 # ===== LOG ===== MAIN_IDX = _id_by_name(_normalize_class_names(CLASS_NAMES, len(CLASS_NAMES)), MAIN_CLASS_NAME) if CLASS_NAMES else None msg = (f"[{epoch}/{num_epochs}] train_loss={avg_train_loss:.4f} " f"val_loss={avg_val_loss:.4f} mIoU={m['miou']:.4f} " f"pixAcc={m['pixel_acc']:.4f} lr={cur_lr:.2e} time={elapsed:.1f}s") if CLASS_NAMES is not None: msg += _format_all_classes(m, CLASS_NAMES, MAIN_IDX) if es_metric is not None: # mostra o valor agregado e os f1s das classes do ES cls_names = _normalize_class_names(CLASS_NAMES, len(CLASS_NAMES)) es_names = [cls_names[i] for i in ES_CLASS_IDS] es_vals_str = ", ".join([f"{cls_names[i]}={m['f1_per_class'][i]:.4f}" for i in ES_CLASS_IDS]) msg += f" || ES[{args.es_agg}:{','.join(es_names)}]={es_metric:.4f} ({es_vals_str})" if args.eq_classes and "veg_union" in m: vu = m["veg_union"] msg += f" || VEG(uniao): F1={vu['f1']:.4f} IoU={vu['iou']:.4f} P={vu['precision']:.4f} R={vu['recall']:.4f}" print(msg) # ===== Checkpoints ===== improved_loss = avg_val_loss < best_val_loss - 1e-6 if improved_loss: best_val_loss = avg_val_loss no_imp_loss = 0 tosave = model.state_dict() payload = { "model": tosave, "optimizer": optimizer.state_dict(), "scaler": scaler.state_dict() if 'scaler' in locals() else None, "epoch": epoch, "best_val_loss": best_val_loss, "best_main_class_f1": best_main_class_f1, "best_es": best_es, } if ema: payload["ema"] = ema.ema.state_dict() torch.save(tosave, os.path.join(save_path, f"{MODEL_NAME}_best.pth")) torch.save(payload, os.path.join(save_path, f"{MODEL_NAME}_best_checkpoint.pth")) print("✅ Novo melhor modelo salvo (val_loss).") else: no_imp_loss += 1 # checkpoint por ES (macro/min/harmonic de erva+cana) if es_metric is not None and (es_metric > best_es + args.delta_es_min): best_es = es_metric no_imp_es = 0 tosave = model.state_dict() payload = { "model": tosave, "optimizer": optimizer.state_dict(), "scaler": scaler.state_dict() if 'scaler' in locals() else None, "epoch": epoch, "best_val_loss": best_val_loss, "best_main_class_f1": best_main_class_f1, "best_es": best_es, "es_agg": args.es_agg, "es_classes": args.es_classes, } if ema: payload["ema"] = ema.ema.state_dict() tag = f"best_es_{args.es_agg}_{args.es_classes.replace(',','_')}" torch.save(tosave, os.path.join(save_path, f"{MODEL_NAME}_{tag}.pth")) torch.save(payload, os.path.join(save_path, f"{MODEL_NAME}_{tag}_checkpoint.pth")) print(f"💾 Checkpoint salvo (melhor ES={best_es:.4f} [{args.es_agg}] em {args.es_classes}).") else: no_imp_es += 1 # (opcional) ainda manter checkpoint por F1 da classe principal: if main_class_f1 is not None: if main_class_f1 > best_main_class_f1 + args.delta_f1_min: best_main_class_f1 = main_class_f1 no_imp_f1 = 0 tosave = model.state_dict() payload = { "model": tosave, "optimizer": optimizer.state_dict(), "scaler": scaler.state_dict() if 'scaler' in locals() else None, "epoch": epoch, "best_val_loss": best_val_loss, "best_main_class_f1": best_main_class_f1, "best_es": best_es, } if ema: payload["ema"] = ema.ema.state_dict() torch.save(tosave, os.path.join(save_path, f"{MODEL_NAME}_best_f1_{MAIN_CLASS_NAME}.pth")) torch.save(payload, os.path.join(save_path, f"{MODEL_NAME}_best_f1_{MAIN_CLASS_NAME}_checkpoint.pth")) print(f"🌿💾 Checkpoint salvo (melhor F1 da {MAIN_CLASS_NAME}).") else: no_imp_f1 += 1 else: no_imp_f1 = 0 # ===== Scheduler ===== if active_sched == "cosine": if no_imp_loss >= args.plateau_patience: active_sched = "plateau" print("🔁 Mudando scheduler: Cosine → ReduceLROnPlateau (platô detectado).") no_imp_loss = 0; no_imp_f1 = 0; no_imp_es = 0 epochs_since_switch = 0 plateau.step(avg_val_loss) for g in optimizer.param_groups: g['lr'] = max(g['lr'] * args.plateau_factor, args.min_lr) else: cosine.step() else: plateau.step(avg_val_loss) # se quiser platear por ES, precisa criar um plateau 'mode=max' epochs_since_switch += 1 # ===== Counters ===== if no_imp_loss > 0 or no_imp_f1 > 0 or no_imp_es > 0: print(f"⏳ Sem melhora — loss: {no_imp_loss}/{args.patience_loss}, ES: {no_imp_es}/{args.patience_es}, {MAIN_CLASS_NAME}: {no_imp_f1}/{args.patience_f1}") # ===== Early stopping (usa ES + loss) ===== if (no_imp_loss >= args.patience_loss and (es_metric is None or no_imp_es >= args.patience_es) and (config["model_to_use"] == "main_class" and no_imp_f1 > args.patience_f1) and (active_sched == "cosine" or epochs_since_switch >= args.grace_after_switch)): print("⏹ Early stopping: loss e ES sem melhora.") break # ===== Plots periódicos ===== if epoch % args.plot_every == 0 or epoch == num_epochs: x_epochs = list(range(1, 1 + len(train_loss_history))) # Loss plt.figure() plt.plot(x_epochs, train_loss_history, marker="o", label="Train Loss") plt.plot(x_epochs, val_loss_history, marker="s", label="Val Loss") plt.xlabel("Época"); plt.ylabel("Loss"); plt.grid(True); plt.legend(); plt.title("Curva de Loss") plt.tight_layout(); plt.savefig(os.path.join(save_path, "loss_curve.png")); plt.close() # LR plt.figure() plt.plot(x_epochs, lr_history, marker=".") plt.xlabel("Época"); plt.ylabel("LR"); plt.grid(True); plt.title("Learning Rate") plt.tight_layout(); plt.savefig(os.path.join(save_path, "lr_curve.png")); plt.close() # mIoU e F1 plt.figure() plt.plot(x_epochs, miou_history, marker="^", label="mIoU") if len(f1_history) == len(miou_history): plt.plot(x_epochs, f1_history, marker="*", label=f"F1 {MAIN_CLASS_NAME}") plt.xlabel("Época"); plt.ylabel("Score"); plt.grid(True); plt.legend(); plt.title(f"mIoU / F1({MAIN_CLASS_NAME})") plt.tight_layout(); plt.savefig(os.path.join(save_path, "metrics_curve.png")); plt.close() def parse_args(): ap = argparse.ArgumentParser() # Treino base ap.add_argument("--epochs", type=int, default=120) ap.add_argument("--batch", type=int, default=16) ap.add_argument("--workers", type=int, default=4) ap.add_argument("--lr", type=float, default=3e-4) ap.add_argument("--min-lr", dest="min_lr", type=float, default=1e-6) ap.add_argument("--weight-decay", dest="weight_decay", type=float, default=1e-4) ap.add_argument("--amp", action="store_true", help="Ativa mixed precision") ap.add_argument("--accum", type=int, default=1, help="Acumulação de gradiente") ap.add_argument("--grad-clip", dest="grad_clip", type=float, default=1.0) ap.add_argument("--seed", type=int, default=42) ap.add_argument("--cudnn-benchmark", action="store_true") ap.add_argument("--freeze-bn", action="store_true") # EMA ap.add_argument("--ema", action="store_true", help="Ativa EMA do modelo para validação/checkpoints") ap.add_argument("--ema-decay", dest="ema_decay", type=float, default=0.999, help="Fator de decaimento da EMA (0..1)") # Pesos de classe (ROI) ap.add_argument("--use-weights", action="store_true") ap.add_argument("--cw-alpha", type=float, default=1.05) ap.add_argument("--cw-min", type=float, default=0.3) ap.add_argument("--cw-max", type=float, default=2.0) ap.add_argument("--cw-bgfloor", type=float, default=0.6) ap.add_argument("--cw-max-per-group", type=int, default=300) ap.add_argument("--cw-warmup", type=int, default=8) ap.add_argument("--cw-erva-min", type=float, default=1.2) ap.add_argument("--cw-cana-max", type=float, default=2.2) ap.add_argument("--cw-chao-min", type=float, default=0.5) # Dice mix ap.add_argument("--dice-max", type=float, default=0.3, help="Fraçao máxima de Dice na loss (0..1)") # Scheduler ap.add_argument("--warmup-epochs", type=int, default=5) ap.add_argument("--plateau-factor", type=float, default=0.5) ap.add_argument("--plateau-patience", type=int, default=6) ap.add_argument("--plateau-cooldown", type=int, default=1) # Early stop ap.add_argument("--patience-loss", type=int, default=12) ap.add_argument("--patience-f1", type=int, default=6) ap.add_argument("--delta-f1-min", type=float, default=0.0015) ap.add_argument("--grace-after-switch", type=int, default=4) # Sampler por nome de classe ap.add_argument("--ensure-per-batch", type=str, default="", help="Ex.: 'erva,cana'") # Tversky ap.add_argument("--tversky-classes", type=str, default="", help="Classes (nomes) separadas por vírgula, ex: 'erva,cana'") ap.add_argument("--tversky-weights", type=str, default="", help="Pesos para cada classe, ex: '0.2,0.15' (valor único faz broadcast)") ap.add_argument("--tversky-alphas", type=str, default="", help="Alphas por classe (FP weight). Ex: '0.72,0.35' (valor único faz broadcast, default=0.5)") ap.add_argument("--tversky-betas", type=str, default="", help="Betas por classe (FN weight). Ex: '0.28,0.65' (valor único faz broadcast, default=0.5)") ap.add_argument("--tversky-in-val", action="store_true", help="Se setado, também soma Tversky no loss de validação (default: não)") # Foco do treinamento ap.add_argument("--es-classes", type=str, default="erva,cana", help="Classes usadas no critério de early stop/checkpoint (ex.: 'erva,cana').") ap.add_argument("--es-agg", type=str, default="macro", choices=["macro", "min", "harmonic"], help="Agregação sobre as classes do ES: macro (média), min (mais rígido), harmonic.") ap.add_argument("--delta-es-min", type=float, default=1e-4, help="Mín. variação para considerar melhora no ES.") ap.add_argument("--patience-es", type=int, default=8, help="Paciência (épocas) sem melhora no ES antes de contar para early stop.") # Classes iguais fase 1 ap.add_argument("--eq-classes", type=str, default="", help="Classes da mesma familia, como vegetacao (ex.: 'erva,cana').") # Outros ap.add_argument("--checkpoint", type=str, default=None) ap.add_argument("--plot-every", dest="plot_every", type=int, default=5) return ap.parse_args() if __name__ == "__main__": args = parse_args() train(args)