agrobot_base/Python/OAK/datasets/_5_train_fastscnn.py

117 lines
4.4 KiB
Python

import json
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
with open("config.json", "r") as f:
config = json.load(f)
MODELO = config["camera"]
MODEL_NAME = config["model_name"]
RESOLUCAO = config["resolucao"]
ROI_INICIO = config["roi_inicio"]
ROI_TAMANHO = config["roi_tamanho"]
save_path = os.path.join(MODELO, "backup", config["modelo"], 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[0], RESOLUCAO[1], 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)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.epochs)
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(start_epoch, 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:
x_epochs = list(range(start_epoch, start_epoch + len(loss_history)))
plt.figure()
plt.plot(x_epochs, 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)