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() def load_class_map(labelmap_path): class_map = {} with open(labelmap_path, 'r') as file: for line in file: # Ignorar comentários e linhas vazias if line.startswith('#') or not line.strip(): continue # Extrair informações parts = line.strip().split(':') if len(parts) >= 2: label, color_str = parts[0], parts[1] # Converter a string de cor em uma tupla de inteiros rgb = tuple(map(int, color_str.split(','))) # Atualizar o class_map class_map[rgb] = len(class_map) return class_map # Transformação personalizada para máscaras def mask_transform(mask): mask = np.array(mask) class_map = load_class_map('dataset/labelmap.txt') class_mask = np.zeros((mask.shape[0], mask.shape[1]), dtype=np.int32) for rgb, idx in class_map.items(): match = (mask == rgb).all(axis=-1) class_mask[match] = idx return torch.tensor(class_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): self.images.append(os.path.join(image_dir, img)) # Combine todos os arquivos de máscara dos diretórios fornecidos for mask_dir in mask_dirs: for mask in os.listdir(mask_dir): self.masks.append(os.path.join(mask_dir, mask.replace('jpg', 'jpeg'))) # Asumindo substituição de extensão 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("RGB") if self.transform: image = self.transform(image) if self.mask_transform: mask = self.mask_transform(mask) return image, mask # Instanciando o Dataset # Altere os caminhos conforme necessário image_dirs = ["dataset/images", "dataset/augmented_images"] mask_dirs = ["dataset/masks", "dataset/augmented_masks"] save_dir = "backup/" model_name = "ruasModel" train_dataset = CustomVOC(image_dirs=image_dirs, mask_dirs=mask_dirs, transform=image_transforms, mask_transform=mask_transform) # DataLoader train_loader = DataLoader(train_dataset, batch_size=4, shuffle=True, 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) 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) 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 for epoch in range(start_epoch, num_epochs): model.train() running_loss = 0.0 for images, masks in train_loader: images = images.to(device) masks = masks.to(device) optimizer.zero_grad() out = model(images) outputs = out loss = criterion(outputs, masks) loss.backward() optimizer.step() running_loss += loss.item() correct = 0 total = 0 # Convertendo as previsões em rótulos preditos _, predicted = torch.max(outputs.data, 1) correct += (predicted == masks).sum().item() total += masks.numel() accuracy = 100 * correct / total # Salvando checkpoints checkpoint_path = f'{save_dir}{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': running_loss, }, checkpoint_path) print(f"Checkpoint salvo em {checkpoint_path}") epoch_losses.append(running_loss / len(train_loader)) epoch_accuracies.append(accuracy) # Atualiza os dados do gráfico ax1.plot(epoch_losses, color=color_loss) ax2.plot(epoch_accuracies, color=color_accuracy) # Redesenha o gráfico fig.canvas.draw() fig.canvas.flush_events() print(f"Epoch {epoch+1}, Loss: {running_loss/len(train_loader)}, Accuracy: {accuracy}") # 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()