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

2224 lines
72 KiB
Python

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
_8_train_multihead.py
Treina SegFormer OAK-FCC-3 com tensor multiespectral CHW conforme
config["input_channels"], incluindo bandas físicas e índices derivados,
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,
input_channel_names: Optional[List[str]] = 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.input_channel_names = list(
input_channel_names or SOURCE_CHANNEL_ORDER[:self.channels]
)
self.strict_channels = bool(strict_channels)
self.resize_hw = resize_hw
self.samples = self._collect_samples()
if not self.samples:
raise RuntimeError(f"Nenhuma amostra encontrada em: {self.root}")
self._validate_tensor_contract()
self._add_sample_class_stats()
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"
tensor_channels = None
if meta_path.exists():
meta = load_json(meta_path)
meta_channels = meta.get("channels")
if isinstance(meta_channels, list) and meta_channels:
tensor_channels = [str(c).strip().upper() for c in meta_channels]
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,
"tensor_channels": tensor_channels,
})
return samples
def _resolve_tensor_channel_indices(self, sample: dict, tensor_c: int) -> List[int]:
"""Mapeia por nome os canais salvos para a entrada pedida no config."""
saved_names = sample.get("tensor_channels")
if saved_names:
if len(saved_names) != tensor_c:
raise RuntimeError(
f"Meta/tensor incompatíveis em {sample['tensor']}: "
f"meta.channels={saved_names}, tensor C={tensor_c}."
)
if len(set(saved_names)) != len(saved_names):
raise RuntimeError(
f"Meta possui canais duplicados em {sample['tensor']}: {saved_names}"
)
missing = [name for name in self.input_channel_names if name not in saved_names]
if missing:
raise RuntimeError(
f"Tensor {sample['tensor']} não possui canais {missing}. "
f"Salvos={saved_names}, pedidos={self.input_channel_names}."
)
return [saved_names.index(name) for name in self.input_channel_names]
# Tensor normalizado sem meta: se C coincide, assume contrato final na
# mesma ordem do config. Mantém compatibilidade com datasets antigos.
if tensor_c == self.channels:
return list(range(self.channels))
# Fallback legado: .npy Raw5 sem informação nominal.
if tensor_c == len(SOURCE_CHANNEL_ORDER):
if all(name in SOURCE_CHANNEL_ORDER for name in self.input_channel_names):
return [SOURCE_CHANNEL_ORDER.index(name) for name in self.input_channel_names]
raise RuntimeError(
f"Não foi possível inferir canais de {sample['tensor']}: tensor C={tensor_c}, "
f"pedidos={self.input_channel_names}, meta.channels ausente."
)
def _validate_tensor_contract(self):
"""Valida todos os headers .npy antes da primeira época."""
contracts = set()
for sample in self.samples:
x = np.load(str(sample["tensor"]), mmap_mode="r")
if x.ndim != 3:
raise RuntimeError(f"Tensor inválido {sample['tensor']}: shape={x.shape}")
indices = self._resolve_tensor_channel_indices(sample, int(x.shape[0]))
if len(indices) != self.channels:
raise RuntimeError(
f"Contrato de canais inválido em {sample['tensor']}: indices={indices}"
)
contracts.add((tuple(sample.get("tensor_channels") or []), int(x.shape[0])))
if self.strict_channels and len(contracts) > 1:
raise RuntimeError(
"Foram encontrados contratos de tensor diferentes no mesmo split: "
f"{sorted(contracts, key=str)}"
)
print(
f"[DATA][CHANNELS] root={self.root} pedidos={self.input_channel_names} "
f"contratos={sorted(contracts, key=str)}"
)
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}")
channel_indices = self._resolve_tensor_channel_indices(s, int(x.shape[0]))
x = x[channel_indices, :, :]
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"],
}
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
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)
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", [])
raw_stat_channels = stats.get("channels", [])
stat_channels = (
[str(c).strip().upper() for c in raw_stat_channels]
if isinstance(raw_stat_channels, list)
else []
)
if len(mean) != len(std):
raise RuntimeError(
f"norm_stats incompatível: {p} | mean={len(mean)}, std={len(std)}"
)
if stat_channels:
if len(stat_channels) != len(mean):
raise RuntimeError(
f"norm_stats incompatível: {p} | channels={len(stat_channels)}, "
f"mean/std={len(mean)}"
)
if len(set(stat_channels)) != len(stat_channels):
raise RuntimeError(
f"norm_stats possui canais duplicados: {p} | {stat_channels}"
)
missing = [name for name in input_channel_names if name not in stat_channels]
if missing:
raise RuntimeError(
f"norm_stats {p} não possui os canais {missing}. "
f"Disponíveis={stat_channels}, pedidos={input_channel_names}."
)
stat_indices = [stat_channels.index(name) for name in input_channel_names]
mean = [mean[i] for i in stat_indices]
std = [std[i] for i in stat_indices]
selected_stat_channels = [stat_channels[i] for i in stat_indices]
elif len(mean) == channels:
# Compatibilidade com norm_stats antigo sem nomes, mas já no
# mesmo contrato/ordem do tensor final.
selected_stat_channels = list(input_channel_names)
elif len(mean) == len(SOURCE_CHANNEL_ORDER) and all(
name in SOURCE_CHANNEL_ORDER for name in input_channel_names
):
stat_indices = [SOURCE_CHANNEL_ORDER.index(name) for name in input_channel_names]
mean = [mean[i] for i in stat_indices]
std = [std[i] for i in stat_indices]
selected_stat_channels = list(input_channel_names)
else:
raise RuntimeError(
f"norm_stats sem nomes incompatível: {p} | mean/std={len(mean)}, "
f"pedidos={input_channel_names}."
)
mean_arr = np.asarray(mean, dtype=np.float64)
std_arr = np.asarray(std, dtype=np.float64)
if not np.isfinite(mean_arr).all() or not np.isfinite(std_arr).all():
raise RuntimeError(f"norm_stats contém NaN/Inf: {p}")
if np.any(std_arr <= 0.0):
raise RuntimeError(f"norm_stats contém std <= 0: {p} | std={std}")
print(f"[NORM] usando stats fixos: {p}")
print(f"[NORM] selected_channels={selected_stat_channels}")
return FixedNormalizer(mean, std).to(device), str(p)
if bool(getattr(args, "allow_missing_norm_stats", False)):
print("[NORM][WARN] norm_stats não encontrado. Usando normalize_per_batch.")
return None, None
raise RuntimeError(
"norm_stats não encontrado. O treino reproduzível exige estatísticas fixas. "
"Use --norm_stats <arquivo> ou --allow-missing-norm-stats conscientemente. "
f"Candidatos: {[str(p) for p in candidates]}"
)
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)
SOURCE_CHANNEL_ORDER = ["R", "G", "B", "RE", "NIR"]
DERIVED_CHANNEL_ORDER = ["NDVI", "NDRE"]
SUPPORTED_INPUT_CHANNELS = SOURCE_CHANNEL_ORDER + DERIVED_CHANNEL_ORDER
def get_input_channel_names(config: dict) -> List[str]:
"""
Define, valida e preserva a ordem dos canais que entram no modelo.
Compatibilidade com configs antigas sem input_channels:
channels=3 -> R,G,B
channels=4 -> R,G,B,RE
channels=5 -> R,G,B,RE,NIR
"""
if "input_channels" in config:
configured = config["input_channels"]
if isinstance(configured, str):
names = [str(c).strip().upper() for c in configured.split(",") if str(c).strip()]
else:
names = [str(c).strip().upper() for c in configured if str(c).strip()]
else:
n = int(config.get("channels", 5))
if n < 1 or n > len(SOURCE_CHANNEL_ORDER):
raise RuntimeError(
"Config antiga sem input_channels só suporta channels entre 1 e 5. "
"Para NDVI/NDRE, declare input_channels explicitamente."
)
names = SOURCE_CHANNEL_ORDER[:n]
if not names:
raise RuntimeError("input_channels não pode ser vazio.")
duplicates = sorted({name for name in names if names.count(name) > 1})
if duplicates:
raise RuntimeError(f"Canais duplicados em input_channels: {duplicates}")
invalid = [c for c in names if c not in SUPPORTED_INPUT_CHANNELS]
if invalid:
raise RuntimeError(
f"Canais inválidos em input_channels: {invalid}. "
f"Suportados: {SUPPORTED_INPUT_CHANNELS}"
)
configured_count = config.get("channels")
if configured_count is not None and int(configured_count) != len(names):
raise RuntimeError(
"config.channels incompatível com config.input_channels: "
f"channels={configured_count}, input_channels={names} ({len(names)} canais)."
)
return names
def get_input_channel_indices(config: dict) -> List[int]:
"""Indices no tensor normalizado v2, que já está na ordem final."""
return list(range(len(get_input_channel_names(config))))
# ============================================================
# Modelo multi-head
# ============================================================
def patch_segformer_encoder_input_channels(
segformer_encoder: nn.Module,
input_channel_names: List[str],
):
"""
Altera o primeiro patch embedding do SegFormer conforme os nomes de canal.
R/G/B recebem o filtro pré-treinado correspondente, mesmo quando a ordem
muda. RE/NIR/NDVI/NDRE recebem a média dos filtros RGB pré-treinados.
"""
names = [str(c).strip().upper() for c in input_channel_names]
in_ch = len(names)
if names == ["R", "G", "B"]:
return segformer_encoder
proj = segformer_encoder.encoder.patch_embeddings[0].proj
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 old_weight.shape[1] < 3:
raise RuntimeError(
f"Patch embedding pré-treinado inesperado: weight={tuple(old_weight.shape)}"
)
pretrained_index = {"R": 0, "G": 1, "B": 2}
mean_rgb = old_weight[:, :3, :, :].mean(dim=1)
for dst_index, name in enumerate(names):
if name in pretrained_index:
src_index = pretrained_index[name]
new_proj.weight[:, dst_index, :, :].copy_(
old_weight[:, src_index, :, :]
)
else:
new_proj.weight[:, dst_index, :, :].copy_(mean_rgb)
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: RGB(3) -> {names} ({in_ch}) | "
"extras=mean_rgb"
)
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,
input_channel_names: List[str],
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)))
input_channel_names = [str(c).strip().upper() for c in input_channel_names]
channels = len(input_channel_names)
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, input_channel_names)
base.config.num_channels = int(channels)
base.config.input_channel_names = list(input_channel_names)
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,
input_channel_names: List[str],
heads_config: Dict[str, dict],
semantic_id2label: Dict[int, str],
semantic_label2id: Dict[str, int],
):
model = MultiHeadSegFormer(
backbone=backbone,
input_channel_names=input_channel_names,
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):
model_cfg = getattr(model, "config", None)
ckpt = {
"epoch": epoch,
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
"best": best,
"model_contract": {
"num_channels": int(getattr(model_cfg, "num_channels", 0) or 0),
"input_channel_names": list(
getattr(model_cfg, "input_channel_names", []) or []
),
},
}
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)
saved_contract = ckpt.get("model_contract", {}) or {}
current_cfg = getattr(model, "config", None)
current_channels = int(getattr(current_cfg, "num_channels", 0) or 0)
current_names = list(getattr(current_cfg, "input_channel_names", []) or [])
saved_channels = int(saved_contract.get("num_channels", 0) or 0)
saved_names = list(saved_contract.get("input_channel_names", []) or [])
if saved_channels and saved_channels != current_channels:
raise RuntimeError(
f"Checkpoint incompatível: salvou {saved_channels} canais, "
f"modelo atual pede {current_channels}."
)
if saved_names and current_names and saved_names != current_names:
raise RuntimeError(
f"Checkpoint incompatível: canais salvos={saved_names}, "
f"canais atuais={current_names}."
)
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(
"--allow-missing-norm-stats",
dest="allow_missing_norm_stats",
action="store_true",
help="Permite normalize_per_batch quando norm_stats não existir. Não recomendado para ensaios comparáveis.",
)
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)
parser.add_argument("--balanced_sampler", action="store_true")
parser.add_argument("--samples_per_epoch", type=int, default=0)
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} tensor_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}")
# stats_source_tag identifica a variante completa. Isso evita colisões
# como RGB+RE e RGB+NIR (ambos possuem quatro canais).
save_dir = Path("backup") / model_family / model_name / stats_source_tag
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"Input names : {input_channel_names}")
print(f"Variant tag : {stats_source_tag}")
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,
input_channel_names=input_channel_names,
strict_channels=True,
resize_hw=resize_hw,
)
ds_val = OakFcc3TensorMultiHeadDataset(
val_root,
heads_config=heads_config,
channels=channels,
input_channel_names=input_channel_names,
strict_channels=True,
resize_hw=resize_hw,
)
print(f"[DATA] train={len(ds_train)} | val={len(ds_val)}")
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}")
dl_train = DataLoader(
ds_train,
batch_size=args.batch,
shuffle=train_shuffle,
sampler=train_sampler,
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)
experiment_stats_path = None
if normalizer is not None:
source_stats = load_json(norm_stats_path) if norm_stats_path else {}
selected_norm_stats = {
"channels": list(input_channel_names),
"mean": normalizer.mean.detach().cpu().view(-1).tolist(),
"std": normalizer.std.detach().cpu().view(-1).tolist(),
"pixels": source_stats.get("pixels"),
"source_norm_stats": norm_stats_path,
"stats_source_tag": stats_source_tag,
}
experiment_stats_path = save_dir / "norm_stats.json"
save_json(experiment_stats_path, selected_norm_stats)
print(f"[NORM] stats do ensaio: {experiment_stats_path}")
model = build_model(
backbone=backbone,
input_channel_names=input_channel_names,
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,
"experiment_norm_stats_path": str(experiment_stats_path) if experiment_stats_path else None,
"class_weight_debug": weight_debug,
"input_channel_names": input_channel_names,
"input_channel_indices": input_channel_indices,
"source_channel_names": config.get("source_channels", SOURCE_CHANNEL_ORDER),
"supported_input_channels": SUPPORTED_INPUT_CHANNELS,
"tensor_contract": "normalized_v2_final_channel_order",
"stats_source_tag": stats_source_tag,
"checkpoint_selection_score": {
"score": "0.50*target_iou + 0.25*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()