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 MODELO = "oak-1" MODEL_NAME = "ervasModel" NUM_CLASSES = 3 RESOLUCAO = (512, 512) # 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() output_stride = 16 save_dir = os.path.join(MODELO, "backup") train_path = os.path.join(MODELO, "dataset", "split", "train") val_path = os.path.join(MODELO, "dataset", "split", "val") # 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(RESOLUCAO), 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 train_dataset = CustomVOC( image_dirs=[os.path.join(train_path, "images")], mask_dirs=[os.path.join(train_path, "masks")], transform=image_transforms, mask_transform=mask_transform ) train_loader = DataLoader(train_dataset, batch_size=5, shuffle=True, num_workers=0) val_dataset = CustomVOC( image_dirs=[os.path.join(val_path, "images")], mask_dirs=[os.path.join(val_path, "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 model = create_deeplabv3plus(NUM_CLASSES, output_stride) criterion = torch.nn.CrossEntropyLoss(ignore_index=255) 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() 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 === if not os.path.exists(save_dir): os.mkdir(save_dir) 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()