2024-02-19 13:11:53 +00:00
|
|
|
import torch
|
|
|
|
|
import torch.optim as optim
|
2024-02-19 13:16:40 +00:00
|
|
|
import network.modeling as models
|
2024-02-19 13:11:53 +00:00
|
|
|
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
|
2024-02-19 13:16:40 +00:00
|
|
|
parser = argparse.ArgumentParser(description='Treinamento do modelo DeepLabV3+')
|
2024-02-19 13:11:53 +00:00
|
|
|
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)
|
|
|
|
|
|
2024-02-19 13:16:40 +00:00
|
|
|
# Definindo o modelo DeepLabV3+
|
|
|
|
|
def create_deeplabv3plus(_num_classes, _output_stride):
|
|
|
|
|
model = models.deeplabv3plus_resnet50(num_classes=_num_classes, output_stride=_output_stride)
|
2024-02-19 13:11:53 +00:00
|
|
|
return model
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# Configuração do modelo, perda e otimizador
|
2024-02-19 13:16:40 +00:00
|
|
|
num_classes = 4
|
|
|
|
|
output_stride = 16
|
|
|
|
|
model = create_deeplabv3plus(num_classes, output_stride)
|
2024-02-19 13:11:53 +00:00
|
|
|
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 = []
|
|
|
|
|
|
2024-02-19 13:16:40 +00:00
|
|
|
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
|
|
|
|
|
|
2024-02-19 13:11:53 +00:00
|
|
|
# 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()
|
|
|
|
|
|
2024-02-19 13:16:40 +00:00
|
|
|
out = model(images)
|
|
|
|
|
|
|
|
|
|
outputs = out
|
2024-02-19 13:11:53 +00:00
|
|
|
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)
|
|
|
|
|
|
2024-02-19 13:16:40 +00:00
|
|
|
# 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()
|
|
|
|
|
|
2024-02-19 13:11:53 +00:00
|
|
|
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}")
|
2024-02-19 13:16:40 +00:00
|
|
|
|
|
|
|
|
plt.ioff() # Desativa o modo interativo
|
|
|
|
|
plt.show()
|