import os import time import argparse import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from fast_scnn import FastSCNN from roi_seg_dataset import ROISegDataset import matplotlib.pyplot as plt # ⚙️ Configurações MODELO = "oak-1" MODEL_NAME = "ervas_full" RESOLUCAO = (384, 384) ROI_INICIO = 0.0 ROI_TAMANHO = 1.0 save_path = os.path.join(MODELO, "backup", "fast_scnn", MODEL_NAME) dataset_path = os.path.join(MODELO, "dataset") split_folder = "train" labelmap_path = os.path.join(dataset_path, "labelmap.txt") batch_size = 8 num_workers = 4 def train(args): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Device: {device}") ds_train = ROISegDataset(os.path.join(dataset_path, "split", split_folder), save_path, ROI_INICIO, ROI_TAMANHO, RESOLUCAO[1], RESOLUCAO[0], labelmap_path) dl_train = DataLoader(ds_train, batch_size=batch_size, shuffle=True, num_workers=num_workers, pin_memory=True) model = FastSCNN(num_classes=len(ds_train.classes)).to(device) criterion = nn.CrossEntropyLoss(ignore_index=ds_train.ignore_id) optimizer = optim.AdamW(model.parameters(), lr=args.lr, weight_decay=1e-4) scaler = torch.cuda.amp.GradScaler(enabled=args.amp) start_epoch = 1 best_loss = 1e9 loss_history = [] if args.checkpoint and os.path.exists(args.checkpoint): print(f"🔁 Carregando modelo salvo: {args.checkpoint}") checkpoint = torch.load(args.checkpoint, map_location=device) if "model" in checkpoint: model.load_state_dict(checkpoint["model"]) optimizer.load_state_dict(checkpoint["optimizer"]) scaler.load_state_dict(checkpoint["scaler"]) start_epoch = checkpoint.get("epoch", 1) + 1 best_loss = checkpoint.get("best_loss", 1e9) else: # Caso seja apenas um .pth com model.state_dict() direto model.load_state_dict(checkpoint) for epoch in range(1, args.epochs + 1): model.train() total_loss = 0 t0 = time.time() for x, y in dl_train: x, y = x.to(device), y.to(device) optimizer.zero_grad() with torch.cuda.amp.autocast(enabled=args.amp): logits = model(x) loss = criterion(logits, y) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() total_loss += loss.item() * x.size(0) avg_loss = total_loss / len(ds_train) loss_history.append(avg_loss) print(f"[{epoch}/{args.epochs}] loss={avg_loss:.4f} time={time.time()-t0:.1f}s") if avg_loss < best_loss: best_loss = avg_loss torch.save(model.state_dict(), os.path.join(save_path, f"{MODEL_NAME}_best.pth")) torch.save({ "model": model.state_dict(), "optimizer": optimizer.state_dict(), "scaler": scaler.state_dict(), "epoch": epoch, "best_loss": best_loss }, os.path.join(save_path, f"{MODEL_NAME}_best_checkpoint.pth")) print("✅ Novo melhor modelo salvo!") # Plot da curva de perda if epoch % 5 == 0 or epoch == args.epochs: plt.figure() plt.plot(range(start_epoch, epoch + 1), loss_history, marker="o", label="Loss de Treinamento") plt.xlabel("Época") plt.ylabel("Loss") plt.grid(True) plt.legend() plt.title("Curva de Loss") plt.tight_layout() plt.savefig(os.path.join(save_path, "loss_curve.png")) plt.close() def parse_args(): ap = argparse.ArgumentParser() ap.add_argument("--epochs", type=int, default=30) ap.add_argument("--lr", type=float, default=3e-4) ap.add_argument("--amp", action="store_true") ap.add_argument("--checkpoint", type=str, default=None, help="Caminho do modelo .pth para continuar o treinamento") return ap.parse_args() if __name__ == "__main__": args = parse_args() os.makedirs(save_path, exist_ok=True) train(args)