agrobot_base/Python/OAK/datasets/oak-d/_5_train.py

235 lines
7.8 KiB
Python
Raw Normal View History

2025-07-17 11:19:42 +00:00
import torch
import torch.optim as optim
#import network.modeling as models
import numpy as np
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms as T
from PIL import Image
import os
import argparse
import matplotlib.pyplot as plt
# Configuração do parsing de argumentos
parser = argparse.ArgumentParser(description='Treinamento do modelo DeepLabV3+')
parser.add_argument('--checkpoint', type=str, help='Caminho para o checkpoint de onde continuar o treinamento', default=None)
parser.add_argument('--n_epoch', type=str, help='Número de épocas para o treinamento', default='10')
args = parser.parse_args()
# Transformação personalizada para máscaras
def mask_transform(mask):
return torch.tensor(np.array(mask), dtype=torch.long)
# Transformações para as imagens
image_transforms = T.Compose([
T.Resize((512, 512)),
T.ToTensor(),
T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
# Dataset personalizado
class CustomVOC(Dataset):
def __init__(self, image_dirs, mask_dirs, transform=None, mask_transform=None):
self.transform = transform
self.mask_transform = mask_transform
self.images = []
self.masks = []
# Combine todos os arquivos de imagem dos diretórios fornecidos
for image_dir in image_dirs:
for img in os.listdir(image_dir):
if not img.endswith(".jpg") and not img.endswith(".jpeg"):
continue
self.images.append(os.path.join(image_dir, img))
mask_name = img.replace(".jpg", ".png").replace(".jpeg", ".png")
self.masks.append(os.path.join(mask_dirs[0], mask_name))
def __len__(self):
return len(self.images)
def __getitem__(self, idx):
img_path = self.images[idx]
mask_path = self.masks[idx] # Já ajustado na inicialização
image = Image.open(img_path).convert("RGB")
mask = Image.open(mask_path).convert("L") # 1 canal
if self.transform:
image = self.transform(image)
if self.mask_transform:
mask = self.mask_transform(mask)
return image, mask
# Instanciando o Dataset
save_dir = "backup/"
model_name = "ruasModel"
train_dataset = CustomVOC(
image_dirs=["dataset/split/train/images"],
mask_dirs=["dataset/split/train/masks"],
transform=image_transforms,
mask_transform=mask_transform
)
train_loader = DataLoader(train_dataset, batch_size=4, shuffle=True, num_workers=0)
val_dataset = CustomVOC(
image_dirs=["dataset/split/val/images"],
mask_dirs=["dataset/split/val/masks"],
transform=image_transforms,
mask_transform=mask_transform
)
val_loader = DataLoader(val_dataset, batch_size=4, shuffle=False, num_workers=0)
# Definindo o modelo DeepLabV3+
def create_deeplabv3plus(_num_classes, _output_stride):
#model = models.deeplabv3plus_resnet50(num_classes=_num_classes, output_stride=_output_stride)
from torchvision.models.segmentation import deeplabv3_resnet50
model = deeplabv3_resnet50(weights=None, num_classes=_num_classes, output_stride=_output_stride)
return model
# Configuração do modelo, perda e otimizador
num_classes = 4
output_stride = 16
model = create_deeplabv3plus(num_classes, output_stride)
criterion = torch.nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)
# Preparando o modelo para treinamento
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# Carregar checkpoint se fornecido
start_epoch = 0 # Iniciar do começo se nenhum checkpoint for fornecido
if args.checkpoint:
checkpoint = torch.load(args.checkpoint)
model.load_state_dict(checkpoint['model_state_dict'])
# Carrega o estado do otimizador
optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
# Garante que todos os tensores no otimizador estejam no dispositivo correto
for state in optimizer.state.values():
for k, v in state.items():
if isinstance(v, torch.Tensor):
state[k] = v.to(device)
start_epoch = checkpoint['epoch']
print(f"Continuando o treinamento do checkpoint {args.checkpoint}, a partir da época {start_epoch+1}")
model.to(device)
# Força os BatchNorms internos (como do ASPP) para modo eval
for m in model.modules():
if isinstance(m, torch.nn.BatchNorm2d):
m.eval()
epoch_losses = []
epoch_accuracies = []
plt.ion() # Ativa o modo interativo
fig, ax1 = plt.subplots()
color_loss = 'tab:red'
ax1.set_xlabel('Época')
ax1.set_ylabel('Loss', color=color_loss)
ax1.tick_params(axis='y', labelcolor=color_loss)
ax2 = ax1.twinx() # Instancia um segundo eixo que compartilha o mesmo eixo x
color_accuracy = 'tab:blue'
ax2.set_ylabel('Accuracy', color=color_accuracy) # Definimos o label do eixo y
ax2.tick_params(axis='y', labelcolor=color_accuracy)
fig.tight_layout() # Ajusta o layout para evitar sobreposições
# Treinamento do modelo
num_epochs = int(args.n_epoch)
# Loop de treinamento, ajustado para continuar de onde parou
best_val_loss = float('inf')
for epoch in range(start_epoch, num_epochs):
print(f"Epoca {epoch}")
# === TREINO ===
model.train()
running_loss = 0.0
correct = 0
total = 0
for images, masks in train_loader:
if images.size(0) == 1:
print("[⚠️] Pulando batch com 1 imagem (evita erro no BatchNorm)")
continue
images = images.to(device)
masks = masks.to(device)
optimizer.zero_grad()
print(f"Batch shape: {images.shape}")
outputs = model(images)
loss = criterion(outputs['out'], masks)
loss.backward()
optimizer.step()
running_loss += loss.item()
_, predicted = torch.max(outputs['out'].data, 1)
correct += (predicted == masks).sum().item()
total += masks.numel()
train_loss = running_loss / len(train_loader)
train_accuracy = 100 * correct / total
# === VALIDAÇÃO ===
model.eval()
val_loss = 0.0
val_correct = 0
val_total = 0
with torch.no_grad():
for images, masks in val_loader:
images = images.to(device)
masks = masks.to(device)
outputs = model(images)
loss = criterion(outputs['out'], masks)
val_loss += loss.item()
_, predicted = torch.max(outputs['out'].data, 1)
val_correct += (predicted == masks).sum().item()
val_total += masks.numel()
val_loss /= len(val_loader)
val_accuracy = 100 * val_correct / val_total
# === Checkpoint da época ===
checkpoint_path = os.path.join(save_dir, f"{model_name}_checkpoint_epoch_{epoch+1}.pth")
torch.save({
'epoch': epoch+1,
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'loss': train_loss,
'val_loss': val_loss,
}, checkpoint_path)
print(f"[✔] Checkpoint salvo: {checkpoint_path}")
# === Melhor modelo (baseado no val_loss) ===
if val_loss < best_val_loss:
best_val_loss = val_loss
best_model_path = os.path.join(save_dir, f"{model_name}_best.pth")
torch.save(model.state_dict(), best_model_path)
print(f"[🌟] Novo melhor modelo salvo: {best_model_path}")
# === Plot dados ===
epoch_losses.append(train_loss)
epoch_accuracies.append(train_accuracy)
ax1.plot(epoch_losses, color='tab:red')
ax2.plot(epoch_accuracies, color='tab:blue')
fig.canvas.draw()
fig.canvas.flush_events()
print(f"[📊] Época {epoch+1} | Train Loss: {train_loss:.4f} | Train Acc: {train_accuracy:.2f}% | Val Loss: {val_loss:.4f} | Val Acc: {val_accuracy:.2f}%")
# Salvar o modelo final
save_path = f'{save_dir}{model_name}_final.pth'
torch.save(model.state_dict(), save_path)
print(f"Modelo salvo em {save_path}")
plt.ioff() # Desativa o modo interativo
plt.show()