112 lines
4.1 KiB
Python
112 lines
4.1 KiB
Python
|
|
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)
|