#!/usr/bin/env python3 # -*- coding: utf-8 -*- import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "max_split_size_mb:128") import csv 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) # ============================================================ # Labelmap # ============================================================ def load_labelmap(labelmap_path: str): """ Tenta usar utils.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 utils.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 # Aceita: # classe # id classe # id,classe,... 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 # ============================================================ # Dataset novo OAK-FCC-3 # ============================================================ class OakFcc3TensorSegDataset(Dataset): """ Lê o contrato pós-normalização/split: root/ group//tensors/.npy # float32 CHW [5,H,W] group//masks/.npy # int HW """ def __init__( self, root: str | Path, channels: int = 5, strict_channels: bool = True, resize_hw: Optional[Tuple[int, int]] = None, ignore_index: int = 255, ): self.root = Path(root) self.channels = int(channels) self.strict_channels = bool(strict_channels) self.resize_hw = resize_hw self.ignore_index = int(ignore_index) self.samples = self._collect_samples() 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" masks_dir = group_dir / "masks" metas_dir = group_dir / "metas" previews_dir = group_dir / "previews" if not tensors_dir.is_dir() or not masks_dir.is_dir(): continue for tensor_path in sorted(tensors_dir.glob("*.npy")): base = tensor_path.stem mask_path = masks_dir / f"{base}.npy" if not mask_path.exists(): print(f"[WARN] Sem mask para tensor: {tensor_path}") 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, "mask": mask_path, "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_mask(self, x: torch.Tensor, y: torch.Tensor): if self.resize_hw is None: return x, y 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) 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 x, y def __getitem__(self, idx): s = self.samples[idx] x = np.load(str(s["tensor"])).astype(np.float32) y = np.load(str(s["mask"])) if x.ndim != 3: raise RuntimeError(f"Tensor inválido {s['tensor']}: shape={x.shape}") if self.strict_channels and x.shape[0] != self.channels: raise RuntimeError( f"Channels inválido em {s['tensor']}: veio {x.shape[0]}, esperado {self.channels}" ) if x.shape[0] > self.channels: x = x[:self.channels] y = y.astype(np.int64) xt = torch.from_numpy(np.ascontiguousarray(x)).float() yt = torch.from_numpy(np.ascontiguousarray(y)).long() xt, yt = self._resize_tensor_mask(xt, yt) return { "image": xt, "mask": yt, "group": s["group"], "base": s["base"], } def collate_fn(batch): imgs = torch.stack([b["image"] for b in batch], dim=0) masks = torch.stack([b["mask"] for b in batch], dim=0) return imgs, masks # ============================================================ # 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): channels = int(config.get("channels", 5)) 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) != channels or len(std) != channels: raise RuntimeError( f"norm_stats incompatível: {p} | channels={channels}, mean={len(mean)}, std={len(std)}" ) print(f"[NORM] usando stats fixos: {p}") print(f"[NORM] channels={stat_channels}") return FixedNormalizer(mean, std).to(device), str(p) print("[NORM] norm_stats não encontrado. Usando normalize_per_batch.") return None, None # ============================================================ # Modelo # ============================================================ def patch_segformer_input_channels(model: 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 model proj = model.segformer.encoder.patch_embeddings[0].proj if proj.in_channels == in_ch: return model 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) model.segformer.encoder.patch_embeddings[0].proj = new_proj model.config.num_channels = in_ch print(f"[MODEL] patch input channels: 3 -> {in_ch}") return model def build_model( backbone: str, num_classes: int, channels: int, id2label: Dict[int, str], label2id: Dict[str, int], ): model = SegformerForSemanticSegmentation.from_pretrained( backbone, num_labels=num_classes, id2label={int(k): str(v) for k, v in id2label.items()}, label2id={str(k): int(v) for k, v in label2id.items()}, ignore_mismatched_sizes=True, ) patch_segformer_input_channels(model, channels) 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 = 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 estimate_class_weights(ds, num_classes, ignore_index=255, max_samples=800, seed=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["mask"].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 # ============================================================ # Train / Val # ============================================================ def run_one_epoch( model, loader, optimizer, device, num_classes, ignore_index, criterion, amp, scaler, train, grad_accum=1, normalizer=None, ): model.train(train) total_loss = 0.0 n_batches = 0 cm = torch.zeros((num_classes, num_classes), 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) in enumerate(loader): imgs = imgs.to(device, non_blocking=True) masks = masks.to(device, non_blocking=True) 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"): out = model(pixel_values=imgs) logits = out.logits if logits.shape[-2:] != masks.shape[-2:]: logits = F.interpolate( logits, size=masks.shape[-2:], mode="bilinear", align_corners=False, ) ce = criterion(logits, masks) dice = dice_loss( logits=logits, target=masks, num_classes=num_classes, ignore_index=ignore_index, smooth=1.0, ) loss = 0.7 * ce + 0.3 * dice 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 with torch.no_grad(): preds = torch.argmax(logits, dim=1) update_confusion_matrix(cm, preds, masks, num_classes, ignore_index) avg_loss = total_loss / max(1, n_batches) miou, iou_per_class = compute_iou_from_cm(cm) acc = compute_pixel_acc_from_cm(cm) return { "loss": avg_loss, "miou": miou, "iou_per_class": iou_per_class, "acc": acc, "time_s": time.time() - t0, } # ============================================================ # Checkpoint / logs # ============================================================ def save_checkpoint(path, model, optimizer, scaler, epoch, best_miou, best_main_iou, extra=None): ckpt = { "epoch": epoch, "model": model.state_dict(), "optimizer": optimizer.state_dict(), "best_miou": best_miou, "best_main_iou": best_main_iou, } 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) # ============================================================ # 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("--main_class", default=None) 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) args = parser.parse_args() set_seed(args.seed) config = load_json(args.config) if config.get("dual_head", False): raise RuntimeError("Este script está preparado para single-head. No config, use dual_head=false.") W, H = config["resolucao"] channels = int(config.get("channels", 5)) backbone = config.get("backbone", "nvidia/mit-b1") fusion_mode = config.get("fusion_mode", "stacked") 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}") main_class_name = str(args.main_class or config.get("main_class_name", "cana")).lower() 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}") id2label, 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) num_classes = len(id2label) if num_classes <= 1: raise RuntimeError(f"num_classes inválido: {num_classes}") main_class_id = label2id.get(main_class_name) print("==========================================") print("Train SegFormer OAK-FCC-3") 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"Classes : {num_classes} -> {id2label}") print(f"Ignore index : {ignore_index}") print(f"Main class : {main_class_name} -> {main_class_id}") 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 = OakFcc3TensorSegDataset( train_root, channels=channels, strict_channels=True, resize_hw=resize_hw, ignore_index=ignore_index, ) ds_val = OakFcc3TensorSegDataset( val_root, channels=channels, strict_channels=True, resize_hw=resize_hw, ignore_index=ignore_index, ) print(f"[DATA] train={len(ds_train)} | val={len(ds_val)}") dl_train = DataLoader( ds_train, batch_size=args.batch, shuffle=True, 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, num_classes=num_classes, channels=channels, id2label=id2label, label2id=label2id, ) if args.grad_ckpt: try: model.gradient_checkpointing_enable() print("[MODEL] gradient checkpointing enabled") except Exception as e: print(f"[WARN] gradient checkpointing não suportado: {e}") model.to(device) if args.class_weights.lower() == "none": weights = None elif args.class_weights.lower() == "auto": w, counts = estimate_class_weights( ds_train, num_classes=num_classes, ignore_index=ignore_index, seed=args.seed, ) weights = w.to(device) print("[LOSS] class counts:", counts.astype(int).tolist()) print("[LOSS] class weights:", w.cpu().numpy().round(3).tolist()) else: parts = [float(x) for x in args.class_weights.split(",")] if len(parts) != num_classes: raise RuntimeError(f"--class_weights precisa ter {num_classes} valores.") weights = torch.tensor(parts, dtype=torch.float32, device=device) criterion = nn.CrossEntropyLoss(weight=weights, ignore_index=ignore_index) 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_miou_path = save_dir / "best_miou.pt" best_main_path = save_dir / "best_main.pt" train_log_path = save_dir / "train_log.csv" save_json(save_dir / "train_config_snapshot.json", { "config": config, "args": vars(args), "id2label": id2label, "label2id": label2id, "ignore_index": ignore_index, "norm_stats_path": norm_stats_path, }) start_epoch = 1 best_miou = -1.0 best_main_iou = -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_miou = float(ckpt.get("best_miou", -1.0)) best_main_iou = float(ckpt.get("best_main_iou", -1.0)) print(f"[RESUME] {resume_path} epoch={start_epoch}") def pretty_iou(iou_list): return " | ".join([ f"{id2label.get(i, i)}:{v:.3f}" for i, v in enumerate(iou_list) ]) 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, num_classes=num_classes, ignore_index=ignore_index, criterion=criterion, amp=args.amp, scaler=scaler, train=True, grad_accum=max(1, args.grad_accum), normalizer=normalizer, ) if device.type == "cuda": torch.cuda.empty_cache() va = run_one_epoch( model=model, loader=dl_val, optimizer=None, device=device, num_classes=num_classes, ignore_index=ignore_index, criterion=criterion, amp=args.amp_val, scaler=None, train=False, grad_accum=1, normalizer=normalizer, ) scheduler.step(va["loss"]) main_iou = va["miou"] if main_class_id is None else va["iou_per_class"][main_class_id] print( f"TRAIN: loss={tr['loss']:.4f} acc={tr['acc']:.4f} " f"miou={tr['miou']:.4f} t={tr['time_s']:.1f}s" ) print( f"VAL : loss={va['loss']:.4f} acc={va['acc']:.4f} " f"miou={va['miou']:.4f} main_iou={float(main_iou):.4f} t={va['time_s']:.1f}s" ) print("IoU:", pretty_iou(va["iou_per_class"])) save_checkpoint( last_path, model, optimizer, scaler=scaler, epoch=epoch, best_miou=best_miou, best_main_iou=best_main_iou, extra={ "train": tr, "val": va, "main_iou": float(main_iou), }, ) 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_miou=best_miou, best_main_iou=best_main_iou, ) improved = False if va["miou"] > best_miou: best_miou = va["miou"] improved = True save_checkpoint( best_miou_path, model, optimizer, scaler=scaler, epoch=epoch, best_miou=best_miou, best_main_iou=best_main_iou, extra={"val": va, "main_iou": float(main_iou)}, ) print(f"[BEST mIoU] {best_miou:.4f} -> {best_miou_path}") if float(main_iou) > best_main_iou: best_main_iou = float(main_iou) improved = True save_checkpoint( best_main_path, model, optimizer, scaler=scaler, epoch=epoch, best_miou=best_miou, best_main_iou=best_main_iou, extra={"val": va, "main_iou": float(main_iou)}, ) print(f"[BEST MAIN] {best_main_iou:.4f} -> {best_main_path}") append_train_log(train_log_path, { "epoch": epoch, "lr": lr_now, "train_loss": tr["loss"], "train_acc": tr["acc"], "train_miou": tr["miou"], "val_loss": va["loss"], "val_acc": va["acc"], "val_miou": va["miou"], "val_main_iou": float(main_iou), "best_miou": best_miou, "best_main_iou": best_main_iou, }) 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 mIoU : {best_miou:.4f}") print(f"Best main IoU : {best_main_iou:.4f}") print(f"Save dir : {save_dir}") if __name__ == "__main__": main()