2026-05-21 22:57:35 +00:00
|
|
|
#!/usr/bin/env python3
|
|
|
|
|
# -*- coding: utf-8 -*-
|
|
|
|
|
|
|
|
|
|
"""
|
|
|
|
|
_8_train_segformer_oak_fcc3_multihead.py
|
|
|
|
|
|
|
|
|
|
Treina SegFormer OAK-FCC-3 com tensor multispectral [R,G,B,RE,NIR]
|
|
|
|
|
e contrato multi-head:
|
|
|
|
|
|
|
|
|
|
1) semantic_head:
|
|
|
|
|
máscara: masks/<base>.npy
|
|
|
|
|
classes: 0=chao, 1=cana, 2=erva, 255=ignore
|
|
|
|
|
|
|
|
|
|
2) vegetation_head:
|
|
|
|
|
máscara: masks_vegetation/<base>.npy
|
|
|
|
|
classes: 0=nao_vegetacao, 1=vegetacao, 255=ignore
|
|
|
|
|
|
|
|
|
|
3) cana_head:
|
|
|
|
|
máscara: masks_cana/<base>.npy
|
|
|
|
|
classes: 0=nao_cana, 1=cana, 255=ignore
|
|
|
|
|
|
|
|
|
|
Decisão operacional futura:
|
|
|
|
|
alvo/pulverizavel = vegetation == 1 AND cana == 0
|
|
|
|
|
|
|
|
|
|
Entrada esperada após normalize + split:
|
|
|
|
|
|
|
|
|
|
dataset/split/train/group/<grupo>/
|
|
|
|
|
tensors/<base>.npy
|
|
|
|
|
masks/<base>.npy
|
|
|
|
|
masks_vegetation/<base>.npy
|
|
|
|
|
masks_cana/<base>.npy
|
|
|
|
|
metas/<base>.json
|
|
|
|
|
previews/<base>.png
|
|
|
|
|
|
|
|
|
|
dataset/split/val/group/<grupo>/...
|
|
|
|
|
|
|
|
|
|
Exemplo:
|
|
|
|
|
|
|
|
|
|
python _8_train_segformer_oak_fcc3_multihead.py ^
|
|
|
|
|
--epochs 50 --batch 2 --lr 3e-5 --wd 0.01 ^
|
|
|
|
|
--num_workers 2 --amp --amp_val --grad_accum 4 ^
|
|
|
|
|
--class_weights auto
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
|
|
|
import os
|
|
|
|
|
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "max_split_size_mb:128")
|
|
|
|
|
|
|
|
|
|
import csv
|
|
|
|
|
import copy
|
|
|
|
|
import json
|
|
|
|
|
import time
|
|
|
|
|
import argparse
|
|
|
|
|
import random
|
|
|
|
|
from pathlib import Path
|
|
|
|
|
from typing import Dict, List, Optional, Tuple, Any
|
|
|
|
|
|
|
|
|
|
import numpy as np
|
|
|
|
|
import torch
|
|
|
|
|
import torch.nn as nn
|
|
|
|
|
import torch.nn.functional as F
|
|
|
|
|
from torch.utils.data import Dataset, DataLoader
|
|
|
|
|
from torch.amp import autocast, GradScaler
|
|
|
|
|
|
|
|
|
|
from transformers import SegformerForSemanticSegmentation
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ============================================================
|
|
|
|
|
# Seed / util
|
|
|
|
|
# ============================================================
|
|
|
|
|
|
|
|
|
|
def set_seed(seed: int = 42):
|
|
|
|
|
random.seed(seed)
|
|
|
|
|
np.random.seed(seed)
|
|
|
|
|
torch.manual_seed(seed)
|
|
|
|
|
torch.cuda.manual_seed_all(seed)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def ensure_dir(path: Path):
|
|
|
|
|
path.mkdir(parents=True, exist_ok=True)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def load_json(path: str | Path) -> dict:
|
|
|
|
|
with open(path, "r", encoding="utf-8") as f:
|
|
|
|
|
return json.load(f)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def save_json(path: Path, data: dict):
|
|
|
|
|
ensure_dir(path.parent)
|
|
|
|
|
with path.open("w", encoding="utf-8") as f:
|
|
|
|
|
json.dump(data, f, ensure_ascii=False, indent=2)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def safe_float(x, default=0.0) -> float:
|
|
|
|
|
try:
|
|
|
|
|
return float(x)
|
|
|
|
|
except Exception:
|
|
|
|
|
return float(default)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ============================================================
|
|
|
|
|
# Labelmap
|
|
|
|
|
# ============================================================
|
|
|
|
|
|
|
|
|
|
def load_labelmap(labelmap_path: str):
|
|
|
|
|
"""
|
|
|
|
|
Tenta usar helpers.carregar_labelmap_completo.
|
|
|
|
|
Fallback simples para labelmap com uma classe por linha.
|
|
|
|
|
"""
|
|
|
|
|
try:
|
|
|
|
|
from helpers import carregar_labelmap_completo, _infer_ignore_id
|
|
|
|
|
|
|
|
|
|
_cor_para_id, _colormap_rgb, id_para_nome, ignore_rgb = carregar_labelmap_completo(labelmap_path)
|
|
|
|
|
ignore_id = _infer_ignore_id(ignore_rgb, 255)
|
|
|
|
|
|
|
|
|
|
id2label = {int(k): str(v) for k, v in id_para_nome.items()}
|
|
|
|
|
label2id = {v.lower(): k for k, v in id2label.items()}
|
|
|
|
|
|
|
|
|
|
return id2label, label2id, int(ignore_id)
|
|
|
|
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
print(f"[WARN] Não consegui usar helpers.carregar_labelmap_completo: {e}")
|
|
|
|
|
print("[WARN] Usando parser simples: uma classe por linha.")
|
|
|
|
|
|
|
|
|
|
id2label = {}
|
|
|
|
|
with open(labelmap_path, "r", encoding="utf-8") as f:
|
|
|
|
|
for line in f:
|
|
|
|
|
s = line.strip()
|
|
|
|
|
if not s or s.startswith("#"):
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
parts = s.replace(",", " ").split()
|
|
|
|
|
|
|
|
|
|
if len(parts) >= 2 and parts[0].isdigit():
|
|
|
|
|
cid = int(parts[0])
|
|
|
|
|
name = parts[1]
|
|
|
|
|
else:
|
|
|
|
|
cid = len(id2label)
|
|
|
|
|
name = parts[0]
|
|
|
|
|
|
|
|
|
|
if name.lower() in ("ignore", "void", "background_ignore"):
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
id2label[cid] = name
|
|
|
|
|
|
|
|
|
|
label2id = {v.lower(): k for k, v in id2label.items()}
|
|
|
|
|
return id2label, label2id, 255
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ============================================================
|
|
|
|
|
# Head config
|
|
|
|
|
# ============================================================
|
|
|
|
|
|
|
|
|
|
DEFAULT_HEADS = {
|
|
|
|
|
"semantic": {
|
|
|
|
|
"enabled": True,
|
|
|
|
|
"type": "multiclass",
|
|
|
|
|
"num_classes": 3,
|
|
|
|
|
"mask_dir": "masks",
|
|
|
|
|
"classes": {
|
|
|
|
|
"chao": 0,
|
|
|
|
|
"cana": 1,
|
|
|
|
|
"erva": 2,
|
|
|
|
|
},
|
|
|
|
|
"ignore_index": 255,
|
|
|
|
|
"loss_weight": 0.10,
|
|
|
|
|
},
|
|
|
|
|
"vegetation": {
|
|
|
|
|
"enabled": True,
|
|
|
|
|
"type": "binary",
|
|
|
|
|
"num_classes": 2,
|
|
|
|
|
"mask_dir": "masks_vegetation",
|
|
|
|
|
"classes": {
|
|
|
|
|
"background": 0,
|
|
|
|
|
"vegetation": 1,
|
|
|
|
|
},
|
|
|
|
|
"ignore_index": 255,
|
|
|
|
|
"loss_weight": 0.20,
|
|
|
|
|
},
|
|
|
|
|
"cana": {
|
|
|
|
|
"enabled": True,
|
|
|
|
|
"type": "binary",
|
|
|
|
|
"num_classes": 2,
|
|
|
|
|
"mask_dir": "masks_cana",
|
|
|
|
|
"classes": {
|
|
|
|
|
"not_cana": 0,
|
|
|
|
|
"cana": 1,
|
|
|
|
|
},
|
|
|
|
|
"ignore_index": 255,
|
|
|
|
|
"loss_weight": 0.25,
|
|
|
|
|
},
|
|
|
|
|
"target": {
|
|
|
|
|
"enabled": True,
|
|
|
|
|
"type": "binary",
|
|
|
|
|
"num_classes": 2,
|
|
|
|
|
"mask_dir": "__derived_target__",
|
|
|
|
|
"classes": {
|
|
|
|
|
"background": 0,
|
|
|
|
|
"target": 1,
|
|
|
|
|
},
|
|
|
|
|
"ignore_index": 255,
|
|
|
|
|
"loss_weight": 0.45,
|
|
|
|
|
"derived_from": ["vegetation", "cana"],
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def merge_dict(dst: dict, src: dict) -> dict:
|
|
|
|
|
out = copy.deepcopy(dst)
|
|
|
|
|
|
|
|
|
|
def rec(a, b):
|
|
|
|
|
for k, v in b.items():
|
|
|
|
|
if isinstance(v, dict) and isinstance(a.get(k), dict):
|
|
|
|
|
rec(a[k], v)
|
|
|
|
|
else:
|
|
|
|
|
a[k] = v
|
|
|
|
|
|
|
|
|
|
if isinstance(src, dict):
|
|
|
|
|
rec(out, src)
|
|
|
|
|
|
|
|
|
|
return out
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def build_heads_config(config: dict, ignore_index: int) -> Dict[str, dict]:
|
|
|
|
|
cfg = merge_dict(DEFAULT_HEADS, config.get("heads", {}) or {})
|
|
|
|
|
|
|
|
|
|
# Compatibilidade com configs antigas.
|
|
|
|
|
if not bool(config.get("multi_head", True)):
|
|
|
|
|
print("[WARN] config.multi_head não está true. Este script vai treinar multi-head mesmo assim.")
|
|
|
|
|
|
|
|
|
|
for name, hcfg in cfg.items():
|
|
|
|
|
hcfg.setdefault("enabled", True)
|
|
|
|
|
hcfg.setdefault("ignore_index", ignore_index)
|
|
|
|
|
hcfg["num_classes"] = int(hcfg.get("num_classes", 2))
|
|
|
|
|
hcfg["loss_weight"] = float(hcfg.get("loss_weight", 1.0))
|
|
|
|
|
hcfg["mask_dir"] = str(hcfg.get("mask_dir", "masks"))
|
|
|
|
|
|
|
|
|
|
active = {k: v for k, v in cfg.items() if bool(v.get("enabled", True))}
|
|
|
|
|
|
|
|
|
|
if "semantic" not in active:
|
|
|
|
|
raise RuntimeError("A head 'semantic' deve estar habilitada neste contrato.")
|
|
|
|
|
if "vegetation" not in active:
|
|
|
|
|
raise RuntimeError("A head 'vegetation' deve estar habilitada neste contrato.")
|
|
|
|
|
if "cana" not in active:
|
|
|
|
|
raise RuntimeError("A head 'cana' deve estar habilitada neste contrato.")
|
|
|
|
|
|
|
|
|
|
total_w = sum(float(v.get("loss_weight", 0.0)) for v in active.values())
|
|
|
|
|
if total_w <= 0:
|
|
|
|
|
raise RuntimeError("Soma de loss_weight das heads deve ser > 0.")
|
|
|
|
|
|
|
|
|
|
# Normaliza os pesos para não bagunçar magnitude da loss.
|
|
|
|
|
for v in active.values():
|
|
|
|
|
v["loss_weight_norm"] = float(v.get("loss_weight", 0.0)) / total_w
|
|
|
|
|
|
|
|
|
|
return active
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ============================================================
|
|
|
|
|
# Dataset OAK-FCC-3 multi-head
|
|
|
|
|
# ============================================================
|
|
|
|
|
|
|
|
|
|
class OakFcc3TensorMultiHeadDataset(Dataset):
|
|
|
|
|
"""
|
|
|
|
|
Lê o contrato pós-normalização/split:
|
|
|
|
|
|
|
|
|
|
root/
|
|
|
|
|
group/<grupo>/tensors/<base>.npy
|
|
|
|
|
group/<grupo>/masks/<base>.npy
|
|
|
|
|
group/<grupo>/masks_vegetation/<base>.npy
|
|
|
|
|
group/<grupo>/masks_cana/<base>.npy
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
def __init__(
|
|
|
|
|
self,
|
|
|
|
|
root: str | Path,
|
|
|
|
|
heads_config: Dict[str, dict],
|
|
|
|
|
channels: int = 5,
|
|
|
|
|
channel_indices: Optional[List[int]] = None,
|
|
|
|
|
strict_channels: bool = False,
|
|
|
|
|
resize_hw: Optional[Tuple[int, int]] = None,
|
|
|
|
|
):
|
|
|
|
|
self.root = Path(root)
|
|
|
|
|
self.heads_config = heads_config
|
|
|
|
|
self.channels = int(channels)
|
|
|
|
|
self.strict_channels = bool(strict_channels)
|
|
|
|
|
self.resize_hw = resize_hw
|
|
|
|
|
self.channel_indices = channel_indices
|
|
|
|
|
|
|
|
|
|
self.samples = self._collect_samples()
|
2026-06-03 10:37:15 +00:00
|
|
|
self._add_sample_class_stats()
|
2026-05-21 22:57:35 +00:00
|
|
|
|
|
|
|
|
if not self.samples:
|
|
|
|
|
raise RuntimeError(f"Nenhuma amostra encontrada em: {self.root}")
|
|
|
|
|
|
|
|
|
|
def _collect_samples(self):
|
|
|
|
|
samples = []
|
|
|
|
|
|
|
|
|
|
group_root = self.root / "group"
|
|
|
|
|
if not group_root.is_dir():
|
|
|
|
|
raise RuntimeError(f"Pasta group não encontrada em: {self.root}")
|
|
|
|
|
|
|
|
|
|
for group_dir in sorted(group_root.iterdir()):
|
|
|
|
|
if not group_dir.is_dir():
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
tensors_dir = group_dir / "tensors"
|
|
|
|
|
metas_dir = group_dir / "metas"
|
|
|
|
|
previews_dir = group_dir / "previews"
|
|
|
|
|
|
|
|
|
|
if not tensors_dir.is_dir():
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
for tensor_path in sorted(tensors_dir.glob("*.npy")):
|
|
|
|
|
base = tensor_path.stem
|
|
|
|
|
|
|
|
|
|
masks = {}
|
|
|
|
|
missing = []
|
|
|
|
|
|
|
|
|
|
for head_name, hcfg in self.heads_config.items():
|
|
|
|
|
mask_dir_name = str(hcfg.get("mask_dir"))
|
|
|
|
|
|
|
|
|
|
# Head target derivada: não precisa arquivo .npy físico
|
|
|
|
|
if mask_dir_name == "__derived_target__" or bool(hcfg.get("derived", False)):
|
|
|
|
|
masks[head_name] = None
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
mask_dir = group_dir / mask_dir_name
|
|
|
|
|
mask_path = mask_dir / f"{base}.npy"
|
|
|
|
|
|
|
|
|
|
if not mask_path.exists():
|
|
|
|
|
missing.append(f"{head_name}:{mask_path}")
|
|
|
|
|
else:
|
|
|
|
|
masks[head_name] = mask_path
|
|
|
|
|
|
|
|
|
|
if missing:
|
|
|
|
|
print(f"[WARN] Pulando {tensor_path}, masks ausentes: {missing}")
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
meta_path = metas_dir / f"{base}.json"
|
|
|
|
|
preview_path = previews_dir / f"{base}.png"
|
|
|
|
|
|
|
|
|
|
samples.append({
|
|
|
|
|
"group": group_dir.name,
|
|
|
|
|
"base": base,
|
|
|
|
|
"tensor": tensor_path,
|
|
|
|
|
"masks": masks,
|
|
|
|
|
"meta": meta_path if meta_path.exists() else None,
|
|
|
|
|
"preview": preview_path if preview_path.exists() else None,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
return samples
|
|
|
|
|
|
|
|
|
|
def __len__(self):
|
|
|
|
|
return len(self.samples)
|
|
|
|
|
|
|
|
|
|
def _resize_tensor(self, x: torch.Tensor):
|
|
|
|
|
if self.resize_hw is None:
|
|
|
|
|
return x
|
|
|
|
|
|
|
|
|
|
h, w = self.resize_hw
|
|
|
|
|
|
|
|
|
|
if x.shape[-2:] != (h, w):
|
|
|
|
|
x = F.interpolate(
|
|
|
|
|
x.unsqueeze(0),
|
|
|
|
|
size=(h, w),
|
|
|
|
|
mode="bilinear",
|
|
|
|
|
align_corners=False,
|
|
|
|
|
).squeeze(0)
|
|
|
|
|
|
|
|
|
|
return x
|
|
|
|
|
|
|
|
|
|
def _resize_mask(self, y: torch.Tensor):
|
|
|
|
|
if self.resize_hw is None:
|
|
|
|
|
return y
|
|
|
|
|
|
|
|
|
|
h, w = self.resize_hw
|
|
|
|
|
|
|
|
|
|
if y.shape[-2:] != (h, w):
|
|
|
|
|
y = F.interpolate(
|
|
|
|
|
y.unsqueeze(0).unsqueeze(0).float(),
|
|
|
|
|
size=(h, w),
|
|
|
|
|
mode="nearest",
|
|
|
|
|
).squeeze(0).squeeze(0).long()
|
|
|
|
|
|
|
|
|
|
return y
|
|
|
|
|
|
|
|
|
|
def __getitem__(self, idx):
|
|
|
|
|
s = self.samples[idx]
|
|
|
|
|
|
|
|
|
|
x = np.load(str(s["tensor"])).astype(np.float32)
|
|
|
|
|
|
|
|
|
|
if x.ndim != 3:
|
|
|
|
|
raise RuntimeError(f"Tensor inválido {s['tensor']}: shape={x.shape}")
|
|
|
|
|
|
|
|
|
|
if self.channel_indices is not None:
|
|
|
|
|
max_idx = max(self.channel_indices)
|
|
|
|
|
if x.shape[0] <= max_idx:
|
|
|
|
|
raise RuntimeError(
|
|
|
|
|
f"Tensor {s['tensor']} tem {x.shape[0]} canais, "
|
|
|
|
|
f"mas precisa acessar índice {max_idx}. Shape={x.shape}"
|
|
|
|
|
)
|
|
|
|
|
x = x[self.channel_indices, :, :]
|
|
|
|
|
else:
|
|
|
|
|
if x.shape[0] < self.channels:
|
|
|
|
|
raise RuntimeError(
|
|
|
|
|
f"Tensor {s['tensor']} tem {x.shape[0]} canais, "
|
|
|
|
|
f"mas config pediu {self.channels}."
|
|
|
|
|
)
|
|
|
|
|
x = x[:self.channels]
|
|
|
|
|
|
|
|
|
|
xt = torch.from_numpy(np.ascontiguousarray(x)).float()
|
|
|
|
|
xt = self._resize_tensor(xt)
|
|
|
|
|
|
|
|
|
|
masks = {}
|
|
|
|
|
|
|
|
|
|
# Primeiro carrega as heads com arquivo físico
|
|
|
|
|
for head_name, path in s["masks"].items():
|
|
|
|
|
hcfg = self.heads_config[head_name]
|
|
|
|
|
|
|
|
|
|
if path is None:
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
y = np.load(str(path)).astype(np.int64)
|
|
|
|
|
yt = torch.from_numpy(np.ascontiguousarray(y)).long()
|
|
|
|
|
yt = self._resize_mask(yt)
|
|
|
|
|
masks[head_name] = yt
|
|
|
|
|
|
|
|
|
|
# Depois deriva a target, se existir no contrato
|
|
|
|
|
if "target" in self.heads_config:
|
|
|
|
|
if "vegetation" not in masks or "cana" not in masks:
|
|
|
|
|
raise RuntimeError("Head target requer masks vegetation e cana carregadas.")
|
|
|
|
|
|
|
|
|
|
ignore_index = int(self.heads_config["target"].get("ignore_index", 255))
|
|
|
|
|
|
|
|
|
|
veg = masks["vegetation"]
|
|
|
|
|
cana = masks["cana"]
|
|
|
|
|
|
|
|
|
|
valid = (veg != ignore_index) & (cana != ignore_index)
|
|
|
|
|
|
|
|
|
|
target = torch.zeros_like(veg, dtype=torch.long)
|
|
|
|
|
target[(veg == 1) & (cana == 0)] = 1
|
|
|
|
|
target[~valid] = ignore_index
|
|
|
|
|
|
|
|
|
|
masks["target"] = target
|
|
|
|
|
|
|
|
|
|
return {
|
|
|
|
|
"image": xt,
|
|
|
|
|
"masks": masks,
|
|
|
|
|
"group": s["group"],
|
|
|
|
|
"base": s["base"],
|
|
|
|
|
}
|
|
|
|
|
|
2026-06-03 10:37:15 +00:00
|
|
|
def _add_sample_class_stats(self):
|
|
|
|
|
for s in self.samples:
|
|
|
|
|
stats = {
|
|
|
|
|
"pixels_total": 0,
|
|
|
|
|
"pixels_chao": 0,
|
|
|
|
|
"pixels_cana": 0,
|
|
|
|
|
"pixels_erva": 0,
|
|
|
|
|
"pixels_vegetation": 0,
|
|
|
|
|
"pixels_target": 0,
|
|
|
|
|
"pct_cana": 0.0,
|
|
|
|
|
"pct_erva": 0.0,
|
|
|
|
|
"pct_target": 0.0,
|
|
|
|
|
"has_cana": False,
|
|
|
|
|
"has_erva": False,
|
|
|
|
|
"has_target": False,
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
sem_path = s["masks"].get("semantic")
|
|
|
|
|
veg_path = s["masks"].get("vegetation")
|
|
|
|
|
cana_path = s["masks"].get("cana")
|
|
|
|
|
|
|
|
|
|
if sem_path is not None and Path(sem_path).exists():
|
|
|
|
|
sem = np.load(str(sem_path)).astype(np.int64)
|
|
|
|
|
valid = sem != 255
|
|
|
|
|
total = int(valid.sum())
|
|
|
|
|
stats["pixels_total"] = total
|
|
|
|
|
|
|
|
|
|
if total > 0:
|
|
|
|
|
stats["pixels_chao"] = int(((sem == 0) & valid).sum())
|
|
|
|
|
stats["pixels_cana"] = int(((sem == 1) & valid).sum())
|
|
|
|
|
stats["pixels_erva"] = int(((sem == 2) & valid).sum())
|
|
|
|
|
stats["pct_cana"] = stats["pixels_cana"] / total
|
|
|
|
|
stats["pct_erva"] = stats["pixels_erva"] / total
|
|
|
|
|
|
|
|
|
|
if veg_path is not None and cana_path is not None and Path(veg_path).exists() and Path(cana_path).exists():
|
|
|
|
|
veg = np.load(str(veg_path)).astype(np.int64)
|
|
|
|
|
cana = np.load(str(cana_path)).astype(np.int64)
|
|
|
|
|
|
|
|
|
|
valid = (veg != 255) & (cana != 255)
|
|
|
|
|
total = int(valid.sum())
|
|
|
|
|
|
|
|
|
|
if total > 0:
|
|
|
|
|
target = (veg == 1) & (cana == 0) & valid
|
|
|
|
|
stats["pixels_vegetation"] = int(((veg == 1) & valid).sum())
|
|
|
|
|
stats["pixels_target"] = int(target.sum())
|
|
|
|
|
stats["pct_target"] = stats["pixels_target"] / total
|
|
|
|
|
|
|
|
|
|
stats["has_cana"] = stats["pixels_cana"] > 0
|
|
|
|
|
stats["has_erva"] = stats["pixels_erva"] > 0
|
|
|
|
|
stats["has_target"] = stats["pixels_target"] > 0
|
|
|
|
|
|
|
|
|
|
s["class_stats"] = stats
|
|
|
|
|
|
2026-05-21 22:57:35 +00:00
|
|
|
|
|
|
|
|
def collate_fn(batch):
|
|
|
|
|
imgs = torch.stack([b["image"] for b in batch], dim=0)
|
|
|
|
|
|
|
|
|
|
head_names = list(batch[0]["masks"].keys())
|
|
|
|
|
masks = {
|
|
|
|
|
h: torch.stack([b["masks"][h] for b in batch], dim=0)
|
|
|
|
|
for h in head_names
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
meta = {
|
|
|
|
|
"group": [b["group"] for b in batch],
|
|
|
|
|
"base": [b["base"] for b in batch],
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
return imgs, masks, meta
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ============================================================
|
|
|
|
|
# Normalização
|
|
|
|
|
# ============================================================
|
|
|
|
|
|
|
|
|
|
class FixedNormalizer(nn.Module):
|
|
|
|
|
def __init__(self, mean: List[float], std: List[float]):
|
|
|
|
|
super().__init__()
|
|
|
|
|
mean_t = torch.tensor(mean, dtype=torch.float32).view(1, -1, 1, 1)
|
|
|
|
|
std_t = torch.tensor(std, dtype=torch.float32).view(1, -1, 1, 1)
|
|
|
|
|
self.register_buffer("mean", mean_t)
|
|
|
|
|
self.register_buffer("std", torch.clamp(std_t, min=1e-6))
|
|
|
|
|
|
|
|
|
|
def forward(self, x):
|
|
|
|
|
return (x - self.mean) / self.std
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def normalize_per_batch(x: torch.Tensor, eps: float = 1e-6):
|
|
|
|
|
mean = x.mean(dim=(0, 2, 3), keepdim=True)
|
|
|
|
|
std = x.std(dim=(0, 2, 3), keepdim=True).clamp_min(eps)
|
|
|
|
|
return (x - mean) / std
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def build_normalizer(config: dict, args, device: torch.device):
|
|
|
|
|
input_channel_names = get_input_channel_names(config)
|
|
|
|
|
input_channel_indices = get_input_channel_indices(config)
|
|
|
|
|
channels = len(input_channel_names)
|
|
|
|
|
model_family = config.get("modelo", "segformer")
|
|
|
|
|
model_name = config.get("model_name", "model")
|
|
|
|
|
stats_source_tag = config.get("stats_source_tag", f"stacked_raw{channels}")
|
|
|
|
|
|
|
|
|
|
candidates = []
|
|
|
|
|
|
|
|
|
|
if args.norm_stats:
|
|
|
|
|
candidates.append(Path(args.norm_stats))
|
|
|
|
|
|
|
|
|
|
candidates.append(Path("dataset") / f"{config['resolucao'][0]}x{config['resolucao'][1]}" / "group" / "norm_stats.json")
|
|
|
|
|
candidates.append(Path("backup") / model_family / model_name / stats_source_tag / "norm_stats.json")
|
|
|
|
|
|
|
|
|
|
for p in candidates:
|
|
|
|
|
if p.exists():
|
|
|
|
|
stats = load_json(p)
|
|
|
|
|
mean = stats.get("mean", [])
|
|
|
|
|
std = stats.get("std", [])
|
|
|
|
|
stat_channels = stats.get("channels", [])
|
|
|
|
|
|
|
|
|
|
if len(mean) < max(input_channel_indices) + 1 or len(std) < max(input_channel_indices) + 1:
|
|
|
|
|
raise RuntimeError(
|
|
|
|
|
f"norm_stats incompatível: {p} | "
|
|
|
|
|
f"precisa índices={input_channel_indices}, "
|
|
|
|
|
f"mean={len(mean)}, std={len(std)}"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
mean = [mean[i] for i in input_channel_indices]
|
|
|
|
|
std = [std[i] for i in input_channel_indices]
|
|
|
|
|
|
|
|
|
|
if stat_channels:
|
|
|
|
|
stat_channels = [stat_channels[i] for i in input_channel_indices]
|
|
|
|
|
|
|
|
|
|
print(f"[NORM] usando stats fixos: {p}")
|
|
|
|
|
print(f"[NORM] selected_channels={stat_channels or input_channel_names}")
|
|
|
|
|
return FixedNormalizer(mean, std).to(device), str(p)
|
|
|
|
|
|
|
|
|
|
print("[NORM] norm_stats não encontrado. Usando normalize_per_batch.")
|
|
|
|
|
return None, None
|
|
|
|
|
|
|
|
|
|
|
2026-06-03 10:37:15 +00:00
|
|
|
def build_sample_weights(ds, mode="target_focus"):
|
|
|
|
|
weights = []
|
|
|
|
|
|
|
|
|
|
for s in ds.samples:
|
|
|
|
|
st = s.get("class_stats", {})
|
|
|
|
|
group = str(s.get("group", "")).lower()
|
|
|
|
|
|
|
|
|
|
pct_cana = float(st.get("pct_cana", 0.0))
|
|
|
|
|
pct_erva = float(st.get("pct_erva", 0.0))
|
|
|
|
|
pct_target = float(st.get("pct_target", 0.0))
|
|
|
|
|
|
|
|
|
|
has_cana = bool(st.get("has_cana", False))
|
|
|
|
|
has_erva = bool(st.get("has_erva", False))
|
|
|
|
|
has_target = bool(st.get("has_target", False))
|
|
|
|
|
|
|
|
|
|
w = 1.0
|
|
|
|
|
|
|
|
|
|
# Reduz chão puro
|
|
|
|
|
if not has_cana and not has_erva and not has_target:
|
|
|
|
|
w *= 0.35
|
|
|
|
|
|
|
|
|
|
# Aumenta cana
|
|
|
|
|
if has_cana:
|
|
|
|
|
w *= 1.25
|
|
|
|
|
|
|
|
|
|
# Aumenta erva/target com força
|
|
|
|
|
if has_erva:
|
|
|
|
|
w *= 3.0
|
|
|
|
|
|
|
|
|
|
if has_target:
|
|
|
|
|
w *= 4.0
|
|
|
|
|
|
|
|
|
|
# Bônus suave por área real de target/erva
|
|
|
|
|
w *= 1.0 + min(5.0, 80.0 * pct_target)
|
|
|
|
|
w *= 1.0 + min(3.0, 50.0 * pct_erva)
|
|
|
|
|
|
|
|
|
|
# Evita pesos absurdos
|
|
|
|
|
w = max(0.05, min(w, 20.0))
|
|
|
|
|
weights.append(w)
|
|
|
|
|
|
|
|
|
|
return torch.tensor(weights, dtype=torch.double)
|
|
|
|
|
|
|
|
|
|
|
2026-05-21 22:57:35 +00:00
|
|
|
DEFAULT_CHANNEL_ORDER = ["R", "G", "B", "RE", "NIR"]
|
|
|
|
|
|
|
|
|
|
def get_input_channel_names(config: dict) -> List[str]:
|
|
|
|
|
"""
|
|
|
|
|
Define quais canais entram no modelo.
|
|
|
|
|
|
|
|
|
|
Padrão:
|
|
|
|
|
channels=3 -> R,G,B
|
|
|
|
|
channels=4 -> R,G,B,RE
|
|
|
|
|
channels=5 -> R,G,B,RE,NIR
|
|
|
|
|
|
|
|
|
|
Futuramente permite:
|
|
|
|
|
"input_channels": ["R", "G", "B", "NIR"]
|
|
|
|
|
"""
|
|
|
|
|
if "input_channels" in config:
|
|
|
|
|
names = [str(c).upper() for c in config["input_channels"]]
|
|
|
|
|
else:
|
|
|
|
|
n = int(config.get("channels", 5))
|
|
|
|
|
names = DEFAULT_CHANNEL_ORDER[:n]
|
|
|
|
|
|
|
|
|
|
invalid = [c for c in names if c not in DEFAULT_CHANNEL_ORDER]
|
|
|
|
|
if invalid:
|
|
|
|
|
raise RuntimeError(f"Canais inválidos em input_channels: {invalid}")
|
|
|
|
|
|
|
|
|
|
return names
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def get_input_channel_indices(config: dict) -> List[int]:
|
|
|
|
|
names = get_input_channel_names(config)
|
|
|
|
|
return [DEFAULT_CHANNEL_ORDER.index(c) for c in names]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ============================================================
|
|
|
|
|
# Modelo multi-head
|
|
|
|
|
# ============================================================
|
|
|
|
|
|
|
|
|
|
def patch_segformer_encoder_input_channels(segformer_encoder: nn.Module, in_ch: int):
|
|
|
|
|
"""
|
|
|
|
|
Altera o primeiro patch embedding do SegFormer para aceitar C canais.
|
|
|
|
|
Inicializa canais extras pela média dos pesos RGB.
|
|
|
|
|
"""
|
|
|
|
|
if in_ch == 3:
|
|
|
|
|
return segformer_encoder
|
|
|
|
|
|
|
|
|
|
proj = segformer_encoder.encoder.patch_embeddings[0].proj
|
|
|
|
|
|
|
|
|
|
if proj.in_channels == in_ch:
|
|
|
|
|
return segformer_encoder
|
|
|
|
|
|
|
|
|
|
old_weight = proj.weight.data.clone()
|
|
|
|
|
old_bias = proj.bias.data.clone() if proj.bias is not None else None
|
|
|
|
|
|
|
|
|
|
new_proj = nn.Conv2d(
|
|
|
|
|
in_channels=in_ch,
|
|
|
|
|
out_channels=proj.out_channels,
|
|
|
|
|
kernel_size=proj.kernel_size,
|
|
|
|
|
stride=proj.stride,
|
|
|
|
|
padding=proj.padding,
|
|
|
|
|
dilation=proj.dilation,
|
|
|
|
|
groups=proj.groups,
|
|
|
|
|
bias=proj.bias is not None,
|
|
|
|
|
padding_mode=proj.padding_mode,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
with torch.no_grad():
|
|
|
|
|
if in_ch <= old_weight.shape[1]:
|
|
|
|
|
new_proj.weight.copy_(old_weight[:, :in_ch, :, :])
|
|
|
|
|
else:
|
|
|
|
|
new_proj.weight[:, :old_weight.shape[1], :, :].copy_(old_weight)
|
|
|
|
|
extra = in_ch - old_weight.shape[1]
|
|
|
|
|
mean_w = old_weight.mean(dim=1, keepdim=True)
|
|
|
|
|
new_proj.weight[:, old_weight.shape[1]:, :, :].copy_(mean_w.repeat(1, extra, 1, 1))
|
|
|
|
|
|
|
|
|
|
if old_bias is not None:
|
|
|
|
|
new_proj.bias.copy_(old_bias)
|
|
|
|
|
|
|
|
|
|
segformer_encoder.encoder.patch_embeddings[0].proj = new_proj
|
|
|
|
|
print(f"[MODEL] patch input channels: 3 -> {in_ch}")
|
|
|
|
|
return segformer_encoder
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def replace_segformer_decode_classifier(decode_head: nn.Module, num_classes: int):
|
|
|
|
|
"""
|
|
|
|
|
Troca o classifier final da decode_head do HuggingFace SegFormer.
|
|
|
|
|
"""
|
|
|
|
|
if not hasattr(decode_head, "classifier"):
|
|
|
|
|
raise RuntimeError("decode_head sem atributo classifier. Estrutura SegFormer inesperada.")
|
|
|
|
|
|
|
|
|
|
old = decode_head.classifier
|
|
|
|
|
|
|
|
|
|
if not isinstance(old, nn.Conv2d):
|
|
|
|
|
raise RuntimeError(f"decode_head.classifier não é Conv2d: {type(old)}")
|
|
|
|
|
|
|
|
|
|
new = nn.Conv2d(
|
|
|
|
|
in_channels=old.in_channels,
|
|
|
|
|
out_channels=int(num_classes),
|
|
|
|
|
kernel_size=old.kernel_size,
|
|
|
|
|
stride=old.stride,
|
|
|
|
|
padding=old.padding,
|
|
|
|
|
dilation=old.dilation,
|
|
|
|
|
groups=old.groups,
|
|
|
|
|
bias=old.bias is not None,
|
|
|
|
|
padding_mode=old.padding_mode,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
nn.init.xavier_uniform_(new.weight)
|
|
|
|
|
if new.bias is not None:
|
|
|
|
|
nn.init.zeros_(new.bias)
|
|
|
|
|
|
|
|
|
|
decode_head.classifier = new
|
|
|
|
|
return decode_head
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class MultiHeadSegFormer(nn.Module):
|
|
|
|
|
def __init__(
|
|
|
|
|
self,
|
|
|
|
|
backbone: str,
|
|
|
|
|
channels: int,
|
|
|
|
|
heads_config: Dict[str, dict],
|
|
|
|
|
semantic_id2label: Dict[int, str],
|
|
|
|
|
semantic_label2id: Dict[str, int],
|
|
|
|
|
):
|
|
|
|
|
super().__init__()
|
|
|
|
|
|
|
|
|
|
semantic_classes = int(heads_config["semantic"].get("num_classes", len(semantic_id2label)))
|
|
|
|
|
|
|
|
|
|
base = SegformerForSemanticSegmentation.from_pretrained(
|
|
|
|
|
backbone,
|
|
|
|
|
num_labels=semantic_classes,
|
|
|
|
|
id2label={int(k): str(v) for k, v in semantic_id2label.items()},
|
|
|
|
|
label2id={str(k): int(v) for k, v in semantic_label2id.items()},
|
|
|
|
|
ignore_mismatched_sizes=True,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
patch_segformer_encoder_input_channels(base.segformer, channels)
|
|
|
|
|
base.config.num_channels = int(channels)
|
|
|
|
|
|
|
|
|
|
self.segformer = base.segformer
|
|
|
|
|
self.heads_config = heads_config
|
|
|
|
|
self.decode_heads = nn.ModuleDict()
|
|
|
|
|
|
|
|
|
|
for head_name, hcfg in heads_config.items():
|
|
|
|
|
h = copy.deepcopy(base.decode_head)
|
|
|
|
|
h = replace_segformer_decode_classifier(h, int(hcfg["num_classes"]))
|
|
|
|
|
self.decode_heads[head_name] = h
|
|
|
|
|
|
|
|
|
|
self.config = base.config
|
|
|
|
|
|
|
|
|
|
def forward(self, pixel_values: torch.Tensor) -> Dict[str, torch.Tensor]:
|
|
|
|
|
outputs = self.segformer(
|
|
|
|
|
pixel_values=pixel_values,
|
|
|
|
|
output_hidden_states=True,
|
|
|
|
|
return_dict=True,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
hidden_states = outputs.hidden_states
|
|
|
|
|
logits = {}
|
|
|
|
|
|
|
|
|
|
for head_name, head in self.decode_heads.items():
|
|
|
|
|
logits[head_name] = head(hidden_states)
|
|
|
|
|
|
|
|
|
|
return logits
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def build_model(
|
|
|
|
|
backbone: str,
|
|
|
|
|
channels: int,
|
|
|
|
|
heads_config: Dict[str, dict],
|
|
|
|
|
semantic_id2label: Dict[int, str],
|
|
|
|
|
semantic_label2id: Dict[str, int],
|
|
|
|
|
):
|
|
|
|
|
model = MultiHeadSegFormer(
|
|
|
|
|
backbone=backbone,
|
|
|
|
|
channels=channels,
|
|
|
|
|
heads_config=heads_config,
|
|
|
|
|
semantic_id2label=semantic_id2label,
|
|
|
|
|
semantic_label2id=semantic_label2id,
|
|
|
|
|
)
|
|
|
|
|
return model
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ============================================================
|
|
|
|
|
# Métricas / loss
|
|
|
|
|
# ============================================================
|
|
|
|
|
|
|
|
|
|
@torch.no_grad()
|
|
|
|
|
def update_confusion_matrix(cm, preds, labels, num_classes, ignore_index=255):
|
|
|
|
|
preds = preds.reshape(-1)
|
|
|
|
|
labels = labels.reshape(-1)
|
|
|
|
|
|
|
|
|
|
valid = labels != ignore_index
|
|
|
|
|
preds = preds[valid]
|
|
|
|
|
labels = labels[valid]
|
|
|
|
|
|
|
|
|
|
valid2 = (labels >= 0) & (labels < num_classes) & (preds >= 0) & (preds < num_classes)
|
|
|
|
|
preds = preds[valid2]
|
|
|
|
|
labels = labels[valid2]
|
|
|
|
|
|
|
|
|
|
if labels.numel() == 0:
|
|
|
|
|
return
|
|
|
|
|
|
|
|
|
|
idx = labels * num_classes + preds
|
|
|
|
|
bins = torch.bincount(idx, minlength=num_classes * num_classes)
|
|
|
|
|
cm += bins.view(num_classes, num_classes)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@torch.no_grad()
|
|
|
|
|
def compute_iou_from_cm(cm, eps=1e-6):
|
|
|
|
|
cm = cm.float()
|
|
|
|
|
tp = torch.diag(cm)
|
|
|
|
|
fp = cm.sum(0) - tp
|
|
|
|
|
fn = cm.sum(1) - tp
|
|
|
|
|
denom = tp + fp + fn + eps
|
|
|
|
|
iou = (tp / denom).cpu().tolist()
|
|
|
|
|
miou = float(np.mean(iou))
|
|
|
|
|
return miou, iou
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@torch.no_grad()
|
|
|
|
|
def compute_pixel_acc_from_cm(cm, eps=1e-6):
|
|
|
|
|
cm = cm.float()
|
|
|
|
|
return float(torch.diag(cm).sum() / (cm.sum() + eps))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def dice_loss(logits, target, num_classes, ignore_index=255, smooth=1.0):
|
|
|
|
|
probs = torch.softmax(logits, dim=1)
|
|
|
|
|
valid = target != ignore_index
|
|
|
|
|
|
|
|
|
|
if valid.sum() == 0:
|
|
|
|
|
return logits.new_tensor(0.0)
|
|
|
|
|
|
|
|
|
|
target_clamped = target.clone()
|
|
|
|
|
target_clamped[~valid] = 0
|
|
|
|
|
target_clamped = target_clamped.long()
|
|
|
|
|
|
|
|
|
|
target_1h = F.one_hot(target_clamped, num_classes=num_classes)
|
|
|
|
|
target_1h = target_1h.permute(0, 3, 1, 2).float()
|
|
|
|
|
|
|
|
|
|
valid_f = valid.unsqueeze(1).float()
|
|
|
|
|
probs = probs * valid_f
|
|
|
|
|
target_1h = target_1h * valid_f
|
|
|
|
|
|
|
|
|
|
dims = (0, 2, 3)
|
|
|
|
|
inter = (probs * target_1h).sum(dims)
|
|
|
|
|
den = probs.sum(dims) + target_1h.sum(dims)
|
|
|
|
|
|
|
|
|
|
dice = (2.0 * inter + smooth) / (den + smooth)
|
|
|
|
|
return 1.0 - dice.mean()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def build_target_teacher_from_heads(
|
|
|
|
|
logits_by_head: Dict[str, torch.Tensor],
|
|
|
|
|
target_shape: Tuple[int, int],
|
|
|
|
|
config: dict,
|
|
|
|
|
erva_id: int = 2,
|
|
|
|
|
):
|
|
|
|
|
cfg = config.get("target_distillation", {}) or {}
|
|
|
|
|
|
|
|
|
|
if not bool(cfg.get("enabled", False)):
|
|
|
|
|
return None
|
|
|
|
|
|
|
|
|
|
required = ("semantic", "vegetation", "cana")
|
|
|
|
|
if any(k not in logits_by_head for k in required):
|
|
|
|
|
return None
|
|
|
|
|
|
|
|
|
|
sem = logits_by_head["semantic"]
|
|
|
|
|
veg = logits_by_head["vegetation"]
|
|
|
|
|
cana = logits_by_head["cana"]
|
|
|
|
|
|
|
|
|
|
if sem.shape[-2:] != target_shape:
|
|
|
|
|
sem = F.interpolate(sem, size=target_shape, mode="bilinear", align_corners=False)
|
|
|
|
|
if veg.shape[-2:] != target_shape:
|
|
|
|
|
veg = F.interpolate(veg, size=target_shape, mode="bilinear", align_corners=False)
|
|
|
|
|
if cana.shape[-2:] != target_shape:
|
|
|
|
|
cana = F.interpolate(cana, size=target_shape, mode="bilinear", align_corners=False)
|
|
|
|
|
|
|
|
|
|
p_sem = torch.softmax(sem, dim=1)
|
|
|
|
|
p_veg = torch.softmax(veg, dim=1)
|
|
|
|
|
p_cana = torch.softmax(cana, dim=1)
|
|
|
|
|
|
|
|
|
|
p_sem_erva = p_sem[:, int(erva_id), :, :]
|
|
|
|
|
p_veg_pos = p_veg[:, 1, :, :]
|
|
|
|
|
p_cana_pos = p_cana[:, 1, :, :]
|
|
|
|
|
|
|
|
|
|
w_sem_erva = float(cfg.get("w_sem_erva", 0.45))
|
|
|
|
|
w_veg_not_cana = float(cfg.get("w_veg_not_cana", 0.35))
|
|
|
|
|
w_veg_suppressed = float(cfg.get("w_veg_suppressed", 0.20))
|
|
|
|
|
power = float(cfg.get("cana_suppression_power", 1.5))
|
|
|
|
|
|
|
|
|
|
not_cana = torch.clamp(1.0 - p_cana_pos, 0.0, 1.0)
|
|
|
|
|
|
|
|
|
|
teacher = (
|
|
|
|
|
w_sem_erva * p_sem_erva +
|
|
|
|
|
w_veg_not_cana * p_veg_pos * not_cana +
|
|
|
|
|
w_veg_suppressed * p_veg_pos * torch.pow(not_cana, power)
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
teacher_min = float(cfg.get("teacher_min", 0.0))
|
|
|
|
|
teacher_max = float(cfg.get("teacher_max", 1.0))
|
|
|
|
|
teacher = torch.clamp(teacher, teacher_min, teacher_max)
|
|
|
|
|
|
|
|
|
|
if bool(cfg.get("detach_teacher", True)):
|
|
|
|
|
teacher = teacher.detach()
|
|
|
|
|
|
|
|
|
|
return teacher
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def estimate_head_class_weights(
|
|
|
|
|
ds: Dataset,
|
|
|
|
|
head_name: str,
|
|
|
|
|
num_classes: int,
|
|
|
|
|
ignore_index: int = 255,
|
|
|
|
|
max_samples: int = 800,
|
|
|
|
|
seed: int = 42,
|
|
|
|
|
):
|
|
|
|
|
rng = np.random.default_rng(seed)
|
|
|
|
|
n = min(len(ds), max_samples)
|
|
|
|
|
idxs = rng.choice(len(ds), size=n, replace=False)
|
|
|
|
|
|
|
|
|
|
counts = np.zeros(num_classes, dtype=np.float64)
|
|
|
|
|
|
|
|
|
|
for i in idxs:
|
|
|
|
|
item = ds[i]
|
|
|
|
|
m = item["masks"][head_name].numpy().reshape(-1)
|
|
|
|
|
m = m[m != ignore_index]
|
|
|
|
|
m = m[(m >= 0) & (m < num_classes)]
|
|
|
|
|
|
|
|
|
|
if m.size > 0:
|
|
|
|
|
counts += np.bincount(m, minlength=num_classes)[:num_classes]
|
|
|
|
|
|
|
|
|
|
freq = counts / (counts.sum() + 1e-12)
|
|
|
|
|
freq = np.clip(freq, 1e-12, 1.0)
|
|
|
|
|
|
|
|
|
|
weights = 1.0 / np.log(1.02 + freq)
|
|
|
|
|
weights = weights / weights.mean()
|
|
|
|
|
|
|
|
|
|
return torch.tensor(weights, dtype=torch.float32), counts
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def build_criterions(
|
|
|
|
|
ds_train: Dataset,
|
|
|
|
|
heads_config: Dict[str, dict],
|
|
|
|
|
class_weights_mode: str,
|
|
|
|
|
device: torch.device,
|
|
|
|
|
seed: int,
|
|
|
|
|
):
|
|
|
|
|
criterions = {}
|
|
|
|
|
weight_debug = {}
|
|
|
|
|
|
|
|
|
|
for head_name, hcfg in heads_config.items():
|
|
|
|
|
num_classes = int(hcfg["num_classes"])
|
|
|
|
|
ignore_index = int(hcfg.get("ignore_index", 255))
|
|
|
|
|
|
|
|
|
|
mode = str(class_weights_mode).lower()
|
|
|
|
|
weights = None
|
|
|
|
|
counts = None
|
|
|
|
|
|
|
|
|
|
if mode == "none":
|
|
|
|
|
weights = None
|
|
|
|
|
|
|
|
|
|
elif mode == "auto":
|
|
|
|
|
w, counts = estimate_head_class_weights(
|
|
|
|
|
ds_train,
|
|
|
|
|
head_name=head_name,
|
|
|
|
|
num_classes=num_classes,
|
|
|
|
|
ignore_index=ignore_index,
|
|
|
|
|
seed=seed,
|
|
|
|
|
)
|
|
|
|
|
weights = w.to(device)
|
|
|
|
|
print(f"[LOSS:{head_name}] class counts:", counts.astype(int).tolist())
|
|
|
|
|
print(f"[LOSS:{head_name}] class weights:", w.cpu().numpy().round(3).tolist())
|
|
|
|
|
|
|
|
|
|
else:
|
|
|
|
|
# Formato opcional:
|
|
|
|
|
# semantic=1,2,3;vegetation=1,2;cana=1,4
|
|
|
|
|
# ou, para compatibilidade, uma lista aplicada só na semantic.
|
|
|
|
|
parsed = parse_class_weights_string(class_weights_mode)
|
|
|
|
|
if head_name in parsed:
|
|
|
|
|
parts = parsed[head_name]
|
|
|
|
|
elif head_name == "semantic" and "__single__" in parsed:
|
|
|
|
|
parts = parsed["__single__"]
|
|
|
|
|
else:
|
|
|
|
|
parts = None
|
|
|
|
|
|
|
|
|
|
if parts is not None:
|
|
|
|
|
if len(parts) != num_classes:
|
|
|
|
|
raise RuntimeError(
|
|
|
|
|
f"Pesos da head {head_name} precisam ter {num_classes} valores. Veio: {parts}"
|
|
|
|
|
)
|
|
|
|
|
weights = torch.tensor(parts, dtype=torch.float32, device=device)
|
|
|
|
|
|
|
|
|
|
criterions[head_name] = nn.CrossEntropyLoss(
|
|
|
|
|
weight=weights,
|
|
|
|
|
ignore_index=ignore_index,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
weight_debug[head_name] = {
|
|
|
|
|
"weights": weights.detach().cpu().tolist() if weights is not None else None,
|
|
|
|
|
"counts": counts.astype(int).tolist() if counts is not None else None,
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
return criterions, weight_debug
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def parse_class_weights_string(s: str) -> Dict[str, List[float]]:
|
|
|
|
|
out = {}
|
|
|
|
|
txt = str(s).strip()
|
|
|
|
|
|
|
|
|
|
if not txt:
|
|
|
|
|
return out
|
|
|
|
|
|
|
|
|
|
if "=" not in txt:
|
|
|
|
|
out["__single__"] = [float(x) for x in txt.split(",") if x.strip()]
|
|
|
|
|
return out
|
|
|
|
|
|
|
|
|
|
for block in txt.split(";"):
|
|
|
|
|
block = block.strip()
|
|
|
|
|
if not block:
|
|
|
|
|
continue
|
|
|
|
|
k, v = block.split("=", 1)
|
|
|
|
|
out[k.strip()] = [float(x) for x in v.split(",") if x.strip()]
|
|
|
|
|
|
|
|
|
|
return out
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@torch.no_grad()
|
|
|
|
|
def update_operational_target_cm(
|
|
|
|
|
cm_target,
|
|
|
|
|
pred_veg: torch.Tensor,
|
|
|
|
|
pred_cana: torch.Tensor,
|
|
|
|
|
gt_veg: torch.Tensor,
|
|
|
|
|
gt_cana: torch.Tensor,
|
|
|
|
|
ignore_index: int = 255,
|
|
|
|
|
):
|
|
|
|
|
"""
|
|
|
|
|
Métrica operacional:
|
|
|
|
|
target/alvo = vegetação viva e não cana
|
|
|
|
|
|
|
|
|
|
Usa heads vegetation e cana.
|
|
|
|
|
Classe 0 = não alvo
|
|
|
|
|
Classe 1 = alvo/pulverizável
|
|
|
|
|
"""
|
|
|
|
|
valid = (gt_veg != ignore_index) & (gt_cana != ignore_index)
|
|
|
|
|
|
|
|
|
|
gt_target = ((gt_veg == 1) & (gt_cana == 0)).long()
|
|
|
|
|
pred_target = ((pred_veg == 1) & (pred_cana == 0)).long()
|
|
|
|
|
|
|
|
|
|
gt_target = gt_target[valid]
|
|
|
|
|
pred_target = pred_target[valid]
|
|
|
|
|
|
|
|
|
|
if gt_target.numel() == 0:
|
|
|
|
|
return
|
|
|
|
|
|
|
|
|
|
idx = gt_target.reshape(-1) * 2 + pred_target.reshape(-1)
|
|
|
|
|
bins = torch.bincount(idx, minlength=4)
|
|
|
|
|
cm_target += bins.view(2, 2)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def compute_losses_for_batch(
|
|
|
|
|
logits_by_head: Dict[str, torch.Tensor],
|
|
|
|
|
masks_by_head: Dict[str, torch.Tensor],
|
|
|
|
|
heads_config: Dict[str, dict],
|
|
|
|
|
criterions: Dict[str, nn.Module],
|
|
|
|
|
dice_weight: float = 0.30,
|
|
|
|
|
config: Optional[dict] = None,
|
|
|
|
|
epoch: Optional[int] = None,
|
|
|
|
|
):
|
|
|
|
|
total = None
|
|
|
|
|
loss_parts = {}
|
|
|
|
|
|
|
|
|
|
cfg = config or {}
|
|
|
|
|
target_distill_cfg = cfg.get("target_distillation", {}) or {}
|
|
|
|
|
target_distill_enabled = bool(target_distill_cfg.get("enabled", False))
|
|
|
|
|
target_distill_ramp = get_target_distill_ramp_factor(cfg, epoch)
|
|
|
|
|
target_distill_enabled = target_distill_enabled and target_distill_ramp > 0.0
|
|
|
|
|
|
|
|
|
|
for head_name, logits in logits_by_head.items():
|
|
|
|
|
target = masks_by_head[head_name]
|
|
|
|
|
hcfg = heads_config[head_name]
|
|
|
|
|
num_classes = int(hcfg["num_classes"])
|
|
|
|
|
ignore_index = int(hcfg.get("ignore_index", 255))
|
|
|
|
|
head_weight = float(hcfg.get("loss_weight_norm", 1.0))
|
|
|
|
|
|
|
|
|
|
if logits.shape[-2:] != target.shape[-2:]:
|
|
|
|
|
logits = F.interpolate(
|
|
|
|
|
logits,
|
|
|
|
|
size=target.shape[-2:],
|
|
|
|
|
mode="bilinear",
|
|
|
|
|
align_corners=False,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
ce = criterions[head_name](logits, target)
|
|
|
|
|
dl = dice_loss(
|
|
|
|
|
logits=logits,
|
|
|
|
|
target=target,
|
|
|
|
|
num_classes=num_classes,
|
|
|
|
|
ignore_index=ignore_index,
|
|
|
|
|
smooth=1.0,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
head_loss = (1.0 - dice_weight) * ce + dice_weight * dl
|
|
|
|
|
|
|
|
|
|
if head_name == "target" and target_distill_enabled:
|
|
|
|
|
teacher = build_target_teacher_from_heads(
|
|
|
|
|
logits_by_head=logits_by_head,
|
|
|
|
|
target_shape=target.shape[-2:],
|
|
|
|
|
config=cfg,
|
|
|
|
|
erva_id=int(
|
|
|
|
|
cfg.get("heads", {})
|
|
|
|
|
.get("semantic", {})
|
|
|
|
|
.get("classes", {})
|
|
|
|
|
.get("erva", 2)
|
|
|
|
|
),
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
if teacher is not None:
|
|
|
|
|
logits_target = logits
|
|
|
|
|
|
|
|
|
|
if logits_target.shape[-2:] != target.shape[-2:]:
|
|
|
|
|
logits_target = F.interpolate(
|
|
|
|
|
logits_target,
|
|
|
|
|
size=target.shape[-2:],
|
|
|
|
|
mode="bilinear",
|
|
|
|
|
align_corners=False,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
logits_binary = logits_target[:, 1, :, :] - logits_target[:, 0, :, :]
|
|
|
|
|
|
|
|
|
|
valid = target != ignore_index
|
|
|
|
|
|
|
|
|
|
if valid.any():
|
|
|
|
|
logits_binary_v = logits_binary[valid]
|
|
|
|
|
teacher_v = teacher[valid].clamp(0.0, 1.0)
|
|
|
|
|
|
|
|
|
|
distill = F.binary_cross_entropy_with_logits(
|
|
|
|
|
logits_binary_v,
|
|
|
|
|
teacher_v,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
hard_weight = float(target_distill_cfg.get("hard_weight", 0.70))
|
|
|
|
|
distill_weight = float(target_distill_cfg.get("distill_weight", 0.30))
|
|
|
|
|
|
|
|
|
|
# Aplica ramp-up só na parte destilada
|
|
|
|
|
distill_weight = distill_weight * float(target_distill_ramp)
|
|
|
|
|
|
|
|
|
|
denom = max(hard_weight + distill_weight, 1e-6)
|
|
|
|
|
hard_weight = hard_weight / denom
|
|
|
|
|
distill_weight = distill_weight / denom
|
|
|
|
|
|
|
|
|
|
head_loss = hard_weight * head_loss + distill_weight * distill
|
|
|
|
|
|
|
|
|
|
loss_parts[f"{head_name}_distill"] = distill.detach()
|
|
|
|
|
loss_parts[f"{head_name}_distill_ramp"] = torch.as_tensor(
|
|
|
|
|
target_distill_ramp,
|
|
|
|
|
device=distill.device,
|
|
|
|
|
dtype=distill.dtype,
|
|
|
|
|
).detach()
|
|
|
|
|
|
|
|
|
|
weighted = head_weight * head_loss
|
|
|
|
|
|
|
|
|
|
total = weighted if total is None else total + weighted
|
|
|
|
|
|
|
|
|
|
loss_parts[f"{head_name}_loss"] = head_loss.detach()
|
|
|
|
|
loss_parts[f"{head_name}_ce"] = ce.detach()
|
|
|
|
|
loss_parts[f"{head_name}_dice"] = dl.detach()
|
|
|
|
|
|
|
|
|
|
return total, loss_parts
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def resize_logits_to_target_if_needed(
|
|
|
|
|
logits: torch.Tensor,
|
|
|
|
|
target: torch.Tensor,
|
|
|
|
|
) -> torch.Tensor:
|
|
|
|
|
"""
|
|
|
|
|
Garante que logits estejam no mesmo HxW da máscara.
|
|
|
|
|
|
|
|
|
|
Usado para métricas, sem depender do efeito colateral da loss.
|
|
|
|
|
"""
|
|
|
|
|
if logits.shape[-2:] != target.shape[-2:]:
|
|
|
|
|
logits = F.interpolate(
|
|
|
|
|
logits,
|
|
|
|
|
size=target.shape[-2:],
|
|
|
|
|
mode="bilinear",
|
|
|
|
|
align_corners=False,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
return logits
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def get_target_distill_ramp_factor(config: dict, epoch: Optional[int]) -> float:
|
|
|
|
|
cfg = (config or {}).get("target_distillation", {}) or {}
|
|
|
|
|
|
|
|
|
|
if not bool(cfg.get("enabled", False)):
|
|
|
|
|
return 0.0
|
|
|
|
|
|
|
|
|
|
if not bool(cfg.get("rampup_enabled", True)):
|
|
|
|
|
return 1.0
|
|
|
|
|
|
|
|
|
|
if epoch is None:
|
|
|
|
|
return 1.0
|
|
|
|
|
|
|
|
|
|
start_epoch = int(cfg.get("start_epoch", 8))
|
|
|
|
|
rampup_epochs = int(cfg.get("rampup_epochs", 12))
|
|
|
|
|
|
|
|
|
|
if epoch < start_epoch:
|
|
|
|
|
return 0.0
|
|
|
|
|
|
|
|
|
|
if rampup_epochs <= 0:
|
|
|
|
|
return 1.0
|
|
|
|
|
|
|
|
|
|
t = (float(epoch) - float(start_epoch) + 1.0) / float(rampup_epochs)
|
|
|
|
|
t = max(0.0, min(1.0, t))
|
|
|
|
|
|
|
|
|
|
# Smoothstep: sobe suave, sem tranco
|
|
|
|
|
return float(t * t * (3.0 - 2.0 * t))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ============================================================
|
|
|
|
|
# Train / Val
|
|
|
|
|
# ============================================================
|
|
|
|
|
|
|
|
|
|
def run_one_epoch(
|
|
|
|
|
model,
|
|
|
|
|
loader,
|
|
|
|
|
optimizer,
|
|
|
|
|
device,
|
|
|
|
|
heads_config,
|
|
|
|
|
criterions,
|
|
|
|
|
amp,
|
|
|
|
|
scaler,
|
|
|
|
|
train,
|
|
|
|
|
grad_accum=1,
|
|
|
|
|
normalizer=None,
|
|
|
|
|
dice_weight=0.30,
|
|
|
|
|
config=None,
|
|
|
|
|
epoch: Optional[int] = None,
|
|
|
|
|
):
|
|
|
|
|
model.train(train)
|
|
|
|
|
|
|
|
|
|
total_loss = 0.0
|
|
|
|
|
n_batches = 0
|
|
|
|
|
|
|
|
|
|
loss_sums = {}
|
|
|
|
|
|
|
|
|
|
cms = {
|
|
|
|
|
head_name: torch.zeros(
|
|
|
|
|
(int(hcfg["num_classes"]), int(hcfg["num_classes"])),
|
|
|
|
|
dtype=torch.int64,
|
|
|
|
|
device=device,
|
|
|
|
|
)
|
|
|
|
|
for head_name, hcfg in heads_config.items()
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
cm_target = torch.zeros((2, 2), dtype=torch.int64, device=device)
|
|
|
|
|
|
|
|
|
|
t0 = time.time()
|
|
|
|
|
|
|
|
|
|
with torch.set_grad_enabled(train):
|
|
|
|
|
if train and optimizer is not None:
|
|
|
|
|
optimizer.zero_grad(set_to_none=True)
|
|
|
|
|
|
|
|
|
|
for step, (imgs, masks, _meta) in enumerate(loader):
|
|
|
|
|
imgs = imgs.to(device, non_blocking=True)
|
|
|
|
|
masks = {k: v.to(device, non_blocking=True) for k, v in masks.items()}
|
|
|
|
|
|
|
|
|
|
if normalizer is not None:
|
|
|
|
|
imgs = normalizer(imgs)
|
|
|
|
|
else:
|
|
|
|
|
imgs = normalize_per_batch(imgs)
|
|
|
|
|
|
|
|
|
|
with autocast(device_type="cuda", enabled=amp and device.type == "cuda"):
|
|
|
|
|
logits_by_head = model(pixel_values=imgs)
|
|
|
|
|
|
|
|
|
|
loss, loss_parts = compute_losses_for_batch(
|
|
|
|
|
logits_by_head=logits_by_head,
|
|
|
|
|
masks_by_head=masks,
|
|
|
|
|
heads_config=heads_config,
|
|
|
|
|
criterions=criterions,
|
|
|
|
|
dice_weight=dice_weight,
|
|
|
|
|
config=config,
|
|
|
|
|
epoch=epoch,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
if train and grad_accum > 1:
|
|
|
|
|
loss = loss / grad_accum
|
|
|
|
|
|
|
|
|
|
if train and optimizer is not None:
|
|
|
|
|
if amp and scaler is not None and device.type == "cuda":
|
|
|
|
|
scaler.scale(loss).backward()
|
|
|
|
|
|
|
|
|
|
if ((step + 1) % grad_accum) == 0 or (step + 1) == len(loader):
|
|
|
|
|
scaler.step(optimizer)
|
|
|
|
|
scaler.update()
|
|
|
|
|
optimizer.zero_grad(set_to_none=True)
|
|
|
|
|
else:
|
|
|
|
|
loss.backward()
|
|
|
|
|
|
|
|
|
|
if ((step + 1) % grad_accum) == 0 or (step + 1) == len(loader):
|
|
|
|
|
optimizer.step()
|
|
|
|
|
optimizer.zero_grad(set_to_none=True)
|
|
|
|
|
|
|
|
|
|
loss_value = float(loss.item()) * (grad_accum if train and grad_accum > 1 else 1.0)
|
|
|
|
|
total_loss += loss_value
|
|
|
|
|
n_batches += 1
|
|
|
|
|
|
|
|
|
|
for k, v in loss_parts.items():
|
|
|
|
|
loss_sums[k] = loss_sums.get(k, 0.0) + float(v.item())
|
|
|
|
|
|
|
|
|
|
with torch.no_grad():
|
|
|
|
|
preds_by_head = {}
|
|
|
|
|
|
|
|
|
|
for head_name, logits in logits_by_head.items():
|
|
|
|
|
if head_name not in masks:
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
target_mask = masks[head_name]
|
|
|
|
|
|
|
|
|
|
logits_metric = resize_logits_to_target_if_needed(
|
|
|
|
|
logits=logits,
|
|
|
|
|
target=target_mask,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
preds = torch.argmax(logits_metric, dim=1)
|
|
|
|
|
preds_by_head[head_name] = preds
|
|
|
|
|
|
|
|
|
|
update_confusion_matrix(
|
|
|
|
|
cms[head_name],
|
|
|
|
|
preds,
|
|
|
|
|
target_mask,
|
|
|
|
|
int(heads_config[head_name]["num_classes"]),
|
|
|
|
|
int(heads_config[head_name].get("ignore_index", 255)),
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# Métrica operacional antiga:
|
|
|
|
|
# target_op = vegetation == 1 AND cana == 0
|
|
|
|
|
# Continua útil para comparar com a nova head target direta.
|
|
|
|
|
if "vegetation" in preds_by_head and "cana" in preds_by_head:
|
|
|
|
|
update_operational_target_cm(
|
|
|
|
|
cm_target=cm_target,
|
|
|
|
|
pred_veg=preds_by_head["vegetation"],
|
|
|
|
|
pred_cana=preds_by_head["cana"],
|
|
|
|
|
gt_veg=masks["vegetation"],
|
|
|
|
|
gt_cana=masks["cana"],
|
|
|
|
|
ignore_index=int(heads_config["vegetation"].get("ignore_index", 255)),
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
avg_loss = total_loss / max(1, n_batches)
|
|
|
|
|
|
|
|
|
|
metrics = {
|
|
|
|
|
"loss": avg_loss,
|
|
|
|
|
"time_s": time.time() - t0,
|
|
|
|
|
"heads": {},
|
|
|
|
|
"loss_parts": {k: v / max(1, n_batches) for k, v in loss_sums.items()},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
for head_name, cm in cms.items():
|
|
|
|
|
miou, iou_per_class = compute_iou_from_cm(cm)
|
|
|
|
|
acc = compute_pixel_acc_from_cm(cm)
|
|
|
|
|
metrics["heads"][head_name] = {
|
|
|
|
|
"miou": miou,
|
|
|
|
|
"iou_per_class": iou_per_class,
|
|
|
|
|
"acc": acc,
|
|
|
|
|
"cm": cm.detach().cpu().tolist(),
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
target_miou, target_iou = compute_iou_from_cm(cm_target)
|
|
|
|
|
target_acc = compute_pixel_acc_from_cm(cm_target)
|
|
|
|
|
|
|
|
|
|
metrics["operational_target"] = {
|
|
|
|
|
"miou": target_miou,
|
|
|
|
|
"iou_background": target_iou[0],
|
|
|
|
|
"iou_target": target_iou[1],
|
|
|
|
|
"acc": target_acc,
|
|
|
|
|
"cm": cm_target.detach().cpu().tolist(),
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
return metrics
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ============================================================
|
|
|
|
|
# Checkpoint / logs
|
|
|
|
|
# ============================================================
|
|
|
|
|
|
|
|
|
|
def save_checkpoint(path, model, optimizer, scaler, epoch, best: dict, extra=None):
|
|
|
|
|
ckpt = {
|
|
|
|
|
"epoch": epoch,
|
|
|
|
|
"model": model.state_dict(),
|
|
|
|
|
"optimizer": optimizer.state_dict(),
|
|
|
|
|
"best": best,
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
if scaler is not None:
|
|
|
|
|
ckpt["scaler"] = scaler.state_dict()
|
|
|
|
|
|
|
|
|
|
if extra:
|
|
|
|
|
ckpt["extra"] = extra
|
|
|
|
|
|
|
|
|
|
torch.save(ckpt, path)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def load_checkpoint(path, model, optimizer=None, scaler=None, map_location="cpu"):
|
|
|
|
|
ckpt = torch.load(path, map_location=map_location, weights_only=False)
|
|
|
|
|
model.load_state_dict(ckpt["model"], strict=True)
|
|
|
|
|
|
|
|
|
|
if optimizer is not None and "optimizer" in ckpt:
|
|
|
|
|
optimizer.load_state_dict(ckpt["optimizer"])
|
|
|
|
|
|
|
|
|
|
if scaler is not None and "scaler" in ckpt:
|
|
|
|
|
scaler.load_state_dict(ckpt["scaler"])
|
|
|
|
|
|
|
|
|
|
return ckpt
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def append_train_log(path: Path, row: dict):
|
|
|
|
|
ensure_dir(path.parent)
|
|
|
|
|
exists = path.exists()
|
|
|
|
|
|
|
|
|
|
with path.open("a", newline="", encoding="utf-8") as f:
|
|
|
|
|
w = csv.DictWriter(f, fieldnames=list(row.keys()))
|
|
|
|
|
if not exists:
|
|
|
|
|
w.writeheader()
|
|
|
|
|
w.writerow(row)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def flatten_epoch_log(epoch: int, lr: float, tr: dict, va: dict, best: dict) -> dict:
|
|
|
|
|
row = {
|
|
|
|
|
"epoch": epoch,
|
|
|
|
|
"lr": lr,
|
|
|
|
|
"train_loss": tr["loss"],
|
|
|
|
|
"val_loss": va["loss"],
|
|
|
|
|
"best_score": best.get("score", -1.0),
|
|
|
|
|
"best_target_iou": best.get("target_iou", -1.0),
|
|
|
|
|
"best_cana_head_iou": best.get("cana_head_iou", -1.0),
|
|
|
|
|
"best_semantic_miou": best.get("semantic_miou", -1.0),
|
|
|
|
|
"best_target_head_iou": best.get("target_head_iou", -1.0),
|
|
|
|
|
"best_operational_target_iou": best.get("operational_target_iou", -1.0),
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
for prefix, obj in (("train", tr), ("val", va)):
|
|
|
|
|
for head_name, hm in obj.get("heads", {}).items():
|
|
|
|
|
row[f"{prefix}_{head_name}_miou"] = hm.get("miou")
|
|
|
|
|
row[f"{prefix}_{head_name}_acc"] = hm.get("acc")
|
|
|
|
|
|
|
|
|
|
ious = hm.get("iou_per_class", []) or []
|
|
|
|
|
for i, v in enumerate(ious):
|
|
|
|
|
row[f"{prefix}_{head_name}_iou_{i}"] = v
|
|
|
|
|
|
|
|
|
|
op = obj.get("operational_target", {}) or {}
|
|
|
|
|
row[f"{prefix}_target_miou"] = op.get("miou")
|
|
|
|
|
row[f"{prefix}_target_iou"] = op.get("iou_target")
|
|
|
|
|
row[f"{prefix}_target_acc"] = op.get("acc")
|
|
|
|
|
|
|
|
|
|
for k, v in obj.get("loss_parts", {}).items():
|
|
|
|
|
row[f"{prefix}_{k}"] = v
|
|
|
|
|
|
|
|
|
|
return row
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ============================================================
|
|
|
|
|
# Pretty print
|
|
|
|
|
# ============================================================
|
|
|
|
|
|
|
|
|
|
def pretty_iou(names: Dict[int, str], iou_list: List[float]):
|
|
|
|
|
return " | ".join([
|
|
|
|
|
f"{names.get(i, i)}:{v:.3f}"
|
|
|
|
|
for i, v in enumerate(iou_list)
|
|
|
|
|
])
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def binary_iou_text(head_name: str, iou_list: List[float]):
|
|
|
|
|
if head_name == "vegetation":
|
|
|
|
|
names = {0: "bg", 1: "veg"}
|
|
|
|
|
elif head_name == "cana":
|
|
|
|
|
names = {0: "not_cana", 1: "cana"}
|
|
|
|
|
elif head_name == "target":
|
|
|
|
|
names = {0: "bg", 1: "target"}
|
|
|
|
|
else:
|
|
|
|
|
names = {0: "0", 1: "1"}
|
|
|
|
|
|
|
|
|
|
return pretty_iou(names, iou_list)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def compute_selection_score(va: dict) -> dict:
|
|
|
|
|
heads = va.get("heads", {})
|
|
|
|
|
op = va.get("operational_target", {})
|
|
|
|
|
|
|
|
|
|
# Preferir target direta, se existir
|
|
|
|
|
target_head_iou = 0.0
|
|
|
|
|
if "target" in heads:
|
|
|
|
|
vals = heads["target"].get("iou_per_class", []) or []
|
|
|
|
|
if len(vals) > 1:
|
|
|
|
|
target_head_iou = safe_float(vals[1], 0.0)
|
|
|
|
|
|
|
|
|
|
operational_target_iou = safe_float(op.get("iou_target"), 0.0)
|
|
|
|
|
|
|
|
|
|
# Se target head existir, ela manda. Se não existir, usa operacional antigo.
|
|
|
|
|
target_iou = target_head_iou if "target" in heads else operational_target_iou
|
|
|
|
|
|
|
|
|
|
cana_iou = 0.0
|
|
|
|
|
if "cana" in heads:
|
|
|
|
|
vals = heads["cana"].get("iou_per_class", []) or []
|
|
|
|
|
if len(vals) > 1:
|
|
|
|
|
cana_iou = safe_float(vals[1], 0.0)
|
|
|
|
|
|
|
|
|
|
veg_miou = safe_float(heads.get("vegetation", {}).get("miou"), 0.0)
|
|
|
|
|
semantic_miou = safe_float(heads.get("semantic", {}).get("miou"), 0.0)
|
|
|
|
|
|
|
|
|
|
score = (
|
|
|
|
|
0.50 * target_iou +
|
|
|
|
|
0.25 * cana_iou +
|
|
|
|
|
0.15 * veg_miou +
|
|
|
|
|
0.10 * semantic_miou
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
return {
|
|
|
|
|
"score": float(score),
|
|
|
|
|
"target_iou": float(target_iou),
|
|
|
|
|
"target_head_iou": float(target_head_iou),
|
|
|
|
|
"operational_target_iou": float(operational_target_iou),
|
|
|
|
|
"cana_head_iou": float(cana_iou),
|
|
|
|
|
"vegetation_miou": float(veg_miou),
|
|
|
|
|
"semantic_miou": float(semantic_miou),
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ============================================================
|
|
|
|
|
# Main
|
|
|
|
|
# ============================================================
|
|
|
|
|
|
|
|
|
|
def main():
|
|
|
|
|
parser = argparse.ArgumentParser()
|
|
|
|
|
|
|
|
|
|
parser.add_argument("--config", default="config.json")
|
|
|
|
|
parser.add_argument("--epochs", type=int, default=80)
|
|
|
|
|
parser.add_argument("--batch", type=int, default=1)
|
|
|
|
|
parser.add_argument("--lr", type=float, default=3e-5)
|
|
|
|
|
parser.add_argument("--wd", type=float, default=0.01)
|
|
|
|
|
parser.add_argument("--num_workers", type=int, default=2)
|
|
|
|
|
|
|
|
|
|
parser.add_argument("--amp", action="store_true")
|
|
|
|
|
parser.add_argument("--amp_val", action="store_true")
|
|
|
|
|
parser.add_argument("--grad_accum", type=int, default=4)
|
|
|
|
|
parser.add_argument("--grad_ckpt", action="store_true")
|
|
|
|
|
|
|
|
|
|
parser.add_argument("--seed", type=int, default=42)
|
|
|
|
|
parser.add_argument("--ignore_index", type=int, default=None)
|
|
|
|
|
|
|
|
|
|
parser.add_argument("--class_weights", default="auto")
|
|
|
|
|
parser.add_argument("--dice_weight", type=float, default=0.30)
|
|
|
|
|
|
|
|
|
|
parser.add_argument("--norm_stats", default=None)
|
|
|
|
|
parser.add_argument("--src-root", default="dataset/split")
|
|
|
|
|
parser.add_argument("--save-every", type=int, default=10)
|
|
|
|
|
|
|
|
|
|
parser.add_argument("--resume", action="store_true")
|
|
|
|
|
parser.add_argument("--resume-ckpt", default=None)
|
|
|
|
|
|
|
|
|
|
parser.add_argument("--early-stop", type=int, default=25)
|
|
|
|
|
|
2026-06-03 10:37:15 +00:00
|
|
|
parser.add_argument("--balanced_sampler", action="store_true")
|
|
|
|
|
parser.add_argument("--samples_per_epoch", type=int, default=0)
|
|
|
|
|
|
2026-05-21 22:57:35 +00:00
|
|
|
args = parser.parse_args()
|
|
|
|
|
|
|
|
|
|
set_seed(args.seed)
|
|
|
|
|
|
|
|
|
|
config = load_json(args.config)
|
|
|
|
|
|
|
|
|
|
W, H = config["resolucao"]
|
|
|
|
|
channels = int(config.get("channels", 5))
|
|
|
|
|
backbone = config.get("backbone", "nvidia/mit-b1")
|
|
|
|
|
fusion_mode = config.get("fusion_mode", "stacked")
|
|
|
|
|
|
|
|
|
|
input_channel_names = get_input_channel_names(config)
|
|
|
|
|
input_channel_indices = get_input_channel_indices(config)
|
|
|
|
|
channels = len(input_channel_names)
|
|
|
|
|
|
|
|
|
|
print(f"Input channels: {input_channel_names} idx={input_channel_indices}")
|
|
|
|
|
|
|
|
|
|
if fusion_mode != "stacked":
|
|
|
|
|
raise RuntimeError("Este script é para fusion_mode='stacked'.")
|
|
|
|
|
|
|
|
|
|
model_family = config.get("modelo", "segformer")
|
|
|
|
|
model_name = config.get("model_name", "test")
|
|
|
|
|
stats_source_tag = config.get("stats_source_tag", f"stacked_raw{channels}")
|
|
|
|
|
|
|
|
|
|
save_dir = Path("backup") / model_family / model_name / f"{fusion_mode}_raw{channels}"
|
|
|
|
|
ensure_dir(save_dir)
|
|
|
|
|
|
|
|
|
|
labelmap_path = Path("dataset") / "labelmap.txt"
|
|
|
|
|
if not labelmap_path.exists():
|
|
|
|
|
raise RuntimeError(f"Labelmap não encontrado: {labelmap_path}")
|
|
|
|
|
|
|
|
|
|
semantic_id2label, semantic_label2id, ignore_from_labelmap = load_labelmap(str(labelmap_path))
|
|
|
|
|
ignore_index = int(args.ignore_index if args.ignore_index is not None else ignore_from_labelmap)
|
|
|
|
|
|
|
|
|
|
heads_config = build_heads_config(config, ignore_index=ignore_index)
|
|
|
|
|
|
|
|
|
|
# Garante que a semantic conhece o num_classes real do labelmap.
|
|
|
|
|
heads_config["semantic"]["num_classes"] = int(len(semantic_id2label))
|
|
|
|
|
heads_config["semantic"]["ignore_index"] = int(ignore_index)
|
|
|
|
|
|
|
|
|
|
print("==========================================")
|
|
|
|
|
print("Train SegFormer OAK-FCC-3 Multi-Head")
|
|
|
|
|
print(f"Backbone : {backbone}")
|
|
|
|
|
print(f"Save dir : {save_dir}")
|
|
|
|
|
print(f"Split root : {args.src_root}")
|
|
|
|
|
print(f"Resolution : {W}x{H}")
|
|
|
|
|
print(f"Channels : {channels}")
|
|
|
|
|
print(f"Semantic : {heads_config['semantic']['num_classes']} -> {semantic_id2label}")
|
|
|
|
|
print(f"Ignore index : {ignore_index}")
|
|
|
|
|
print("Heads:")
|
|
|
|
|
for name, hcfg in heads_config.items():
|
|
|
|
|
print(
|
|
|
|
|
f" - {name}: classes={hcfg['num_classes']} "
|
|
|
|
|
f"mask_dir={hcfg['mask_dir']} "
|
|
|
|
|
f"loss_weight={hcfg['loss_weight']} norm={hcfg['loss_weight_norm']:.3f}"
|
|
|
|
|
)
|
|
|
|
|
print("==========================================")
|
|
|
|
|
|
|
|
|
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
|
|
|
print(f"Device: {device}")
|
|
|
|
|
|
|
|
|
|
train_root = Path(args.src_root) / "train"
|
|
|
|
|
val_root = Path(args.src_root) / "val"
|
|
|
|
|
|
|
|
|
|
resize_hw = (H, W)
|
|
|
|
|
|
|
|
|
|
ds_train = OakFcc3TensorMultiHeadDataset(
|
|
|
|
|
train_root,
|
|
|
|
|
heads_config=heads_config,
|
|
|
|
|
channels=channels,
|
|
|
|
|
channel_indices=input_channel_indices,
|
|
|
|
|
strict_channels=False,
|
|
|
|
|
resize_hw=resize_hw,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
ds_val = OakFcc3TensorMultiHeadDataset(
|
|
|
|
|
val_root,
|
|
|
|
|
heads_config=heads_config,
|
|
|
|
|
channels=channels,
|
|
|
|
|
channel_indices=input_channel_indices,
|
|
|
|
|
strict_channels=False,
|
|
|
|
|
resize_hw=resize_hw,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
print(f"[DATA] train={len(ds_train)} | val={len(ds_val)}")
|
|
|
|
|
|
2026-06-03 10:37:15 +00:00
|
|
|
train_sampler = None
|
|
|
|
|
train_shuffle = True
|
|
|
|
|
|
|
|
|
|
if args.balanced_sampler:
|
|
|
|
|
from torch.utils.data import WeightedRandomSampler
|
|
|
|
|
sample_weights = build_sample_weights(ds_train)
|
|
|
|
|
num_samples = int(args.samples_per_epoch) if args.samples_per_epoch > 0 else len(ds_train)
|
|
|
|
|
|
|
|
|
|
train_sampler = WeightedRandomSampler(
|
|
|
|
|
weights=sample_weights,
|
|
|
|
|
num_samples=num_samples,
|
|
|
|
|
replacement=True,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
train_shuffle = False
|
|
|
|
|
|
|
|
|
|
print("[SAMPLER] WeightedRandomSampler ativado")
|
|
|
|
|
print(f"[SAMPLER] peso min={float(sample_weights.min()):.3f} "
|
|
|
|
|
f"max={float(sample_weights.max()):.3f} "
|
|
|
|
|
f"mean={float(sample_weights.mean()):.3f}")
|
|
|
|
|
|
2026-05-21 22:57:35 +00:00
|
|
|
dl_train = DataLoader(
|
|
|
|
|
ds_train,
|
|
|
|
|
batch_size=args.batch,
|
2026-06-03 10:37:15 +00:00
|
|
|
shuffle=train_shuffle,
|
|
|
|
|
sampler=train_sampler,
|
2026-05-21 22:57:35 +00:00
|
|
|
num_workers=args.num_workers,
|
|
|
|
|
pin_memory=True,
|
|
|
|
|
collate_fn=collate_fn,
|
|
|
|
|
drop_last=True if len(ds_train) >= args.batch else False,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
dl_val = DataLoader(
|
|
|
|
|
ds_val,
|
|
|
|
|
batch_size=1,
|
|
|
|
|
shuffle=False,
|
|
|
|
|
num_workers=max(0, args.num_workers // 2),
|
|
|
|
|
pin_memory=True,
|
|
|
|
|
collate_fn=collate_fn,
|
|
|
|
|
drop_last=False,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
normalizer, norm_stats_path = build_normalizer(config, args, device)
|
|
|
|
|
|
|
|
|
|
model = build_model(
|
|
|
|
|
backbone=backbone,
|
|
|
|
|
channels=channels,
|
|
|
|
|
heads_config=heads_config,
|
|
|
|
|
semantic_id2label=semantic_id2label,
|
|
|
|
|
semantic_label2id=semantic_label2id,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
if args.grad_ckpt:
|
|
|
|
|
try:
|
|
|
|
|
model.segformer.gradient_checkpointing_enable()
|
|
|
|
|
print("[MODEL] gradient checkpointing enabled")
|
|
|
|
|
except Exception as e:
|
|
|
|
|
print(f"[WARN] gradient checkpointing não suportado: {e}")
|
|
|
|
|
|
|
|
|
|
model.to(device)
|
|
|
|
|
|
|
|
|
|
criterions, weight_debug = build_criterions(
|
|
|
|
|
ds_train=ds_train,
|
|
|
|
|
heads_config=heads_config,
|
|
|
|
|
class_weights_mode=args.class_weights,
|
|
|
|
|
device=device,
|
|
|
|
|
seed=args.seed,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
optimizer = torch.optim.AdamW(
|
|
|
|
|
model.parameters(),
|
|
|
|
|
lr=args.lr,
|
|
|
|
|
weight_decay=args.wd,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
|
|
|
|
|
optimizer,
|
|
|
|
|
mode="min",
|
|
|
|
|
factor=0.5,
|
|
|
|
|
patience=6,
|
|
|
|
|
threshold=1e-4,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
scaler = GradScaler(enabled=args.amp and device.type == "cuda")
|
|
|
|
|
|
|
|
|
|
last_path = save_dir / "last.pt"
|
|
|
|
|
best_score_path = save_dir / "best_score.pt"
|
|
|
|
|
best_target_path = save_dir / "best_target.pt"
|
|
|
|
|
best_cana_path = save_dir / "best_cana_head.pt"
|
|
|
|
|
best_semantic_path = save_dir / "best_semantic_miou.pt"
|
|
|
|
|
train_log_path = save_dir / "train_log.csv"
|
|
|
|
|
|
|
|
|
|
save_json(save_dir / "train_config_snapshot.json", {
|
|
|
|
|
"config": config,
|
|
|
|
|
"args": vars(args),
|
|
|
|
|
"semantic_id2label": semantic_id2label,
|
|
|
|
|
"semantic_label2id": semantic_label2id,
|
|
|
|
|
"heads_config": heads_config,
|
|
|
|
|
"ignore_index": ignore_index,
|
|
|
|
|
"norm_stats_path": norm_stats_path,
|
|
|
|
|
"class_weight_debug": weight_debug,
|
|
|
|
|
"input_channel_names": input_channel_names,
|
|
|
|
|
"input_channel_indices": input_channel_indices,
|
|
|
|
|
"checkpoint_selection_score": {
|
|
|
|
|
"score": "0.45*target_iou + 0.30*cana_head_iou + 0.15*vegetation_miou + 0.10*semantic_miou"
|
|
|
|
|
},
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
start_epoch = 1
|
|
|
|
|
best = {
|
|
|
|
|
"score": -1.0,
|
|
|
|
|
"target_iou": -1.0,
|
|
|
|
|
"target_head_iou": -1.0,
|
|
|
|
|
"operational_target_iou": -1.0,
|
|
|
|
|
"cana_head_iou": -1.0,
|
|
|
|
|
"semantic_miou": -1.0,
|
|
|
|
|
"vegetation_miou": -1.0,
|
|
|
|
|
}
|
|
|
|
|
epochs_without_improve = 0
|
|
|
|
|
|
|
|
|
|
resume_path = Path(args.resume_ckpt) if args.resume_ckpt else last_path
|
|
|
|
|
|
|
|
|
|
if args.resume and resume_path.exists():
|
|
|
|
|
ckpt = load_checkpoint(
|
|
|
|
|
resume_path,
|
|
|
|
|
model,
|
|
|
|
|
optimizer,
|
|
|
|
|
scaler=scaler,
|
|
|
|
|
map_location="cpu",
|
|
|
|
|
)
|
|
|
|
|
start_epoch = int(ckpt["epoch"]) + 1
|
|
|
|
|
best = ckpt.get("best", best)
|
|
|
|
|
print(f"[RESUME] {resume_path} epoch={start_epoch}")
|
|
|
|
|
|
|
|
|
|
for epoch in range(start_epoch, args.epochs + 1):
|
|
|
|
|
lr_now = optimizer.param_groups[0]["lr"]
|
|
|
|
|
print(f"\n==== Epoch {epoch}/{args.epochs} | lr={lr_now:.2e} ====")
|
|
|
|
|
|
|
|
|
|
if device.type == "cuda":
|
|
|
|
|
torch.cuda.empty_cache()
|
|
|
|
|
|
|
|
|
|
tr = run_one_epoch(
|
|
|
|
|
model=model,
|
|
|
|
|
loader=dl_train,
|
|
|
|
|
optimizer=optimizer,
|
|
|
|
|
device=device,
|
|
|
|
|
heads_config=heads_config,
|
|
|
|
|
criterions=criterions,
|
|
|
|
|
amp=args.amp,
|
|
|
|
|
scaler=scaler,
|
|
|
|
|
train=True,
|
|
|
|
|
grad_accum=max(1, args.grad_accum),
|
|
|
|
|
normalizer=normalizer,
|
|
|
|
|
dice_weight=args.dice_weight,
|
|
|
|
|
config=config,
|
|
|
|
|
epoch=epoch,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
if device.type == "cuda":
|
|
|
|
|
torch.cuda.empty_cache()
|
|
|
|
|
|
|
|
|
|
va = run_one_epoch(
|
|
|
|
|
model=model,
|
|
|
|
|
loader=dl_val,
|
|
|
|
|
optimizer=None,
|
|
|
|
|
device=device,
|
|
|
|
|
heads_config=heads_config,
|
|
|
|
|
criterions=criterions,
|
|
|
|
|
amp=args.amp_val,
|
|
|
|
|
scaler=None,
|
|
|
|
|
train=False,
|
|
|
|
|
grad_accum=1,
|
|
|
|
|
normalizer=normalizer,
|
|
|
|
|
dice_weight=args.dice_weight,
|
|
|
|
|
config=config,
|
|
|
|
|
epoch=epoch,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
scheduler.step(va["loss"])
|
|
|
|
|
|
|
|
|
|
score_now = compute_selection_score(va)
|
|
|
|
|
|
|
|
|
|
print(
|
|
|
|
|
f"TRAIN: loss={tr['loss']:.4f} "
|
|
|
|
|
f"sem_mIoU={tr['heads']['semantic']['miou']:.4f} "
|
|
|
|
|
f"veg_mIoU={tr['heads']['vegetation']['miou']:.4f} "
|
|
|
|
|
f"cana_mIoU={tr['heads']['cana']['miou']:.4f} "
|
|
|
|
|
f"targetIoU={tr['operational_target']['iou_target']:.4f} "
|
|
|
|
|
f"t={tr['time_s']:.1f}s"
|
|
|
|
|
)
|
|
|
|
|
target_head_iou_val = 0.0
|
|
|
|
|
if "target" in va["heads"]:
|
|
|
|
|
vals = va["heads"]["target"].get("iou_per_class", []) or []
|
|
|
|
|
if len(vals) > 1:
|
|
|
|
|
target_head_iou_val = vals[1]
|
|
|
|
|
print(
|
|
|
|
|
f"VAL : loss={va['loss']:.4f} "
|
|
|
|
|
f"sem_mIoU={va['heads']['semantic']['miou']:.4f} "
|
|
|
|
|
f"veg_mIoU={va['heads']['vegetation']['miou']:.4f} "
|
|
|
|
|
f"cana_mIoU={va['heads']['cana']['miou']:.4f} "
|
|
|
|
|
f"targetHeadIoU={target_head_iou_val:.4f} "
|
|
|
|
|
f"opTargetIoU={va['operational_target']['iou_target']:.4f} "
|
|
|
|
|
f"score={score_now['score']:.4f} "
|
|
|
|
|
f"t={va['time_s']:.1f}s"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
print("IoU semantic:", pretty_iou(semantic_id2label, va["heads"]["semantic"]["iou_per_class"]))
|
|
|
|
|
print("IoU vegetation:", binary_iou_text("vegetation", va["heads"]["vegetation"]["iou_per_class"]))
|
|
|
|
|
print("IoU cana_head:", binary_iou_text("cana", va["heads"]["cana"]["iou_per_class"]))
|
|
|
|
|
print("IoU target:", f"bg:{va['operational_target']['iou_background']:.3f} | alvo:{va['operational_target']['iou_target']:.3f}")
|
|
|
|
|
if "target" in va["heads"]:
|
|
|
|
|
print("IoU target_head:", binary_iou_text("target", va["heads"]["target"]["iou_per_class"]))
|
|
|
|
|
|
|
|
|
|
save_checkpoint(
|
|
|
|
|
last_path,
|
|
|
|
|
model,
|
|
|
|
|
optimizer,
|
|
|
|
|
scaler=scaler,
|
|
|
|
|
epoch=epoch,
|
|
|
|
|
best=best,
|
|
|
|
|
extra={
|
|
|
|
|
"train": tr,
|
|
|
|
|
"val": va,
|
|
|
|
|
"score_now": score_now,
|
|
|
|
|
},
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
if args.save_every > 0 and epoch % args.save_every == 0:
|
|
|
|
|
save_checkpoint(
|
|
|
|
|
save_dir / f"epoch_{epoch:04d}.pt",
|
|
|
|
|
model,
|
|
|
|
|
optimizer,
|
|
|
|
|
scaler=scaler,
|
|
|
|
|
epoch=epoch,
|
|
|
|
|
best=best,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
improved = False
|
|
|
|
|
|
|
|
|
|
if score_now["score"] > best.get("score", -1.0):
|
|
|
|
|
best["score"] = score_now["score"]
|
|
|
|
|
improved = True
|
|
|
|
|
save_checkpoint(
|
|
|
|
|
best_score_path,
|
|
|
|
|
model,
|
|
|
|
|
optimizer,
|
|
|
|
|
scaler=scaler,
|
|
|
|
|
epoch=epoch,
|
|
|
|
|
best=best,
|
|
|
|
|
extra={"val": va, "score_now": score_now},
|
|
|
|
|
)
|
|
|
|
|
print(f"[BEST SCORE] {best['score']:.4f} -> {best_score_path}")
|
|
|
|
|
|
|
|
|
|
if score_now["target_iou"] > best.get("target_iou", -1.0):
|
|
|
|
|
best["target_iou"] = score_now["target_iou"]
|
|
|
|
|
improved = True
|
|
|
|
|
save_checkpoint(
|
|
|
|
|
best_target_path,
|
|
|
|
|
model,
|
|
|
|
|
optimizer,
|
|
|
|
|
scaler=scaler,
|
|
|
|
|
epoch=epoch,
|
|
|
|
|
best=best,
|
|
|
|
|
extra={"val": va, "score_now": score_now},
|
|
|
|
|
)
|
|
|
|
|
print(f"[BEST TARGET] {best['target_iou']:.4f} -> {best_target_path}")
|
|
|
|
|
|
|
|
|
|
if score_now["cana_head_iou"] > best.get("cana_head_iou", -1.0):
|
|
|
|
|
best["cana_head_iou"] = score_now["cana_head_iou"]
|
|
|
|
|
improved = True
|
|
|
|
|
save_checkpoint(
|
|
|
|
|
best_cana_path,
|
|
|
|
|
model,
|
|
|
|
|
optimizer,
|
|
|
|
|
scaler=scaler,
|
|
|
|
|
epoch=epoch,
|
|
|
|
|
best=best,
|
|
|
|
|
extra={"val": va, "score_now": score_now},
|
|
|
|
|
)
|
|
|
|
|
print(f"[BEST CANA HEAD] {best['cana_head_iou']:.4f} -> {best_cana_path}")
|
|
|
|
|
|
|
|
|
|
if score_now["semantic_miou"] > best.get("semantic_miou", -1.0):
|
|
|
|
|
best["semantic_miou"] = score_now["semantic_miou"]
|
|
|
|
|
improved = True
|
|
|
|
|
save_checkpoint(
|
|
|
|
|
best_semantic_path,
|
|
|
|
|
model,
|
|
|
|
|
optimizer,
|
|
|
|
|
scaler=scaler,
|
|
|
|
|
epoch=epoch,
|
|
|
|
|
best=best,
|
|
|
|
|
extra={"val": va, "score_now": score_now},
|
|
|
|
|
)
|
|
|
|
|
print(f"[BEST SEMANTIC] {best['semantic_miou']:.4f} -> {best_semantic_path}")
|
|
|
|
|
|
|
|
|
|
append_train_log(
|
|
|
|
|
train_log_path,
|
|
|
|
|
flatten_epoch_log(
|
|
|
|
|
epoch=epoch,
|
|
|
|
|
lr=lr_now,
|
|
|
|
|
tr=tr,
|
|
|
|
|
va=va,
|
|
|
|
|
best=best,
|
|
|
|
|
),
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
if improved:
|
|
|
|
|
epochs_without_improve = 0
|
|
|
|
|
else:
|
|
|
|
|
epochs_without_improve += 1
|
|
|
|
|
|
|
|
|
|
if args.early_stop > 0 and epochs_without_improve >= args.early_stop:
|
|
|
|
|
print(f"[EARLY STOP] {epochs_without_improve} épocas sem melhora.")
|
|
|
|
|
break
|
|
|
|
|
|
|
|
|
|
print("\nTreino finalizado.")
|
|
|
|
|
print(f"Best score : {best.get('score', -1.0):.4f}")
|
|
|
|
|
print(f"Best target IoU : {best.get('target_iou', -1.0):.4f}")
|
|
|
|
|
print(f"Best cana-head IoU: {best.get('cana_head_iou', -1.0):.4f}")
|
|
|
|
|
print(f"Best semantic mIoU: {best.get('semantic_miou', -1.0):.4f}")
|
|
|
|
|
print(f"Save dir : {save_dir}")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
if __name__ == "__main__":
|
|
|
|
|
main()
|