agrobot_base/Python/OAK/datasets/_8_train_fastscnn_v2.py

1242 lines
61 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# 👉 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)