171 lines
6.0 KiB
Python
171 lines
6.0 KiB
Python
|
|
import torch
|
||
|
|
from torch.utils.data import Dataset, DataLoader
|
||
|
|
from torchvision import transforms as T
|
||
|
|
from PIL import Image
|
||
|
|
import os
|
||
|
|
import numpy as np
|
||
|
|
import matplotlib.pyplot as plt
|
||
|
|
import network.modeling as models
|
||
|
|
|
||
|
|
# Configurações iniciais
|
||
|
|
NUM_CLASSES = 4
|
||
|
|
OUTPUT_STRIDE = 16
|
||
|
|
MODEL_PATH = 'backup/ruasModel_final.pth'
|
||
|
|
TEST_IMAGE_DIR = "dataset/val_images"
|
||
|
|
TEST_MASK_DIR = "dataset/val_masks"
|
||
|
|
|
||
|
|
# 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
|
||
|
|
|
||
|
|
|
||
|
|
# Transformações para as imagens e máscaras
|
||
|
|
test_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]),
|
||
|
|
])
|
||
|
|
|
||
|
|
def invert_class_map(class_map):
|
||
|
|
"""Inverte o mapeamento de RGB para índice de classe para índice de classe para RGB."""
|
||
|
|
inverted_map = {v: k for k, v in class_map.items()}
|
||
|
|
return inverted_map
|
||
|
|
|
||
|
|
def mask_to_color(mask, class_colors):
|
||
|
|
"""Converte uma máscara de uma única camada em uma imagem RGB, utilizando o mapa de cores fornecido."""
|
||
|
|
# Inicializa uma imagem em branco com 3 canais para RGB
|
||
|
|
color_mask = np.zeros((mask.shape[0], mask.shape[1], 3), dtype=np.uint8)
|
||
|
|
|
||
|
|
for class_value, color in class_colors.items():
|
||
|
|
# Encontra onde a máscara tem o valor da classe atual e define a cor correspondente
|
||
|
|
color_mask[mask == class_value] = color
|
||
|
|
|
||
|
|
return color_mask
|
||
|
|
|
||
|
|
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
|
||
|
|
|
||
|
|
class_map = load_class_map('dataset/labelmap.txt')
|
||
|
|
inverted_class_map = invert_class_map(class_map)
|
||
|
|
|
||
|
|
# Transformação personalizada para máscaras
|
||
|
|
def mask_transform(mask):
|
||
|
|
mask = np.array(mask)
|
||
|
|
print(class_map)
|
||
|
|
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)
|
||
|
|
|
||
|
|
# Carregar o dataset de teste
|
||
|
|
test_dataset = CustomVOC(image_dirs=[TEST_IMAGE_DIR],
|
||
|
|
mask_dirs=[TEST_MASK_DIR],
|
||
|
|
transform=test_transforms,
|
||
|
|
mask_transform=mask_transform) # Ajuste conforme necessário
|
||
|
|
|
||
|
|
test_loader = DataLoader(test_dataset, batch_size=1, shuffle=True, num_workers=0)
|
||
|
|
|
||
|
|
# Carregar o modelo
|
||
|
|
def load_model(model_path):
|
||
|
|
model = models.deeplabv3plus_resnet50(num_classes=NUM_CLASSES, output_stride=OUTPUT_STRIDE)
|
||
|
|
model.load_state_dict(torch.load(model_path))
|
||
|
|
model.eval()
|
||
|
|
return model
|
||
|
|
|
||
|
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||
|
|
|
||
|
|
model = load_model(MODEL_PATH)
|
||
|
|
model.to(device)
|
||
|
|
|
||
|
|
# Função de plotagem atualizada
|
||
|
|
def plot_images(image, true_mask, pred_mask):
|
||
|
|
true_mask_color = mask_to_color(true_mask.cpu().numpy(), inverted_class_map)
|
||
|
|
pred_mask_color = mask_to_color(pred_mask.cpu().numpy(), inverted_class_map)
|
||
|
|
|
||
|
|
fig, axs = plt.subplots(1, 3, figsize=(20, 5))
|
||
|
|
# Adiciona .cpu() antes de .numpy() para mover o tensor para a CPU
|
||
|
|
axs[0].imshow(image.cpu().numpy().transpose(1, 2, 0))
|
||
|
|
axs[0].set_title('Imagem Original')
|
||
|
|
axs[1].imshow(true_mask_color)
|
||
|
|
axs[1].set_title('Máscara Verdadeira')
|
||
|
|
axs[2].imshow(pred_mask_color)
|
||
|
|
axs[2].set_title('Máscara Predita')
|
||
|
|
for ax in axs:
|
||
|
|
ax.axis('off')
|
||
|
|
plt.show()
|
||
|
|
|
||
|
|
|
||
|
|
accuracies = []
|
||
|
|
|
||
|
|
|
||
|
|
# Executando a validação e plotando resultados
|
||
|
|
with torch.no_grad():
|
||
|
|
for image, true_mask in test_loader:
|
||
|
|
image, true_mask = image.to(device), true_mask.to(device)
|
||
|
|
output = model(image)
|
||
|
|
pred_mask = torch.argmax(output.squeeze(), dim=0)
|
||
|
|
#plot_images(image.squeeze(), true_mask.squeeze(), pred_mask)
|
||
|
|
correct_predictions = (pred_mask == true_mask).float() # Converte para float para calcular a média
|
||
|
|
accuracy = correct_predictions.mean() # Calcula a acurácia para a imagem atual
|
||
|
|
accuracies.append(accuracy.item()) # Adiciona a acurácia da imagem atual à lista
|
||
|
|
|
||
|
|
average_accuracy = np.mean(accuracies)
|
||
|
|
print(f"Acurácia média no conjunto de teste: {average_accuracy * 100:.2f}%")
|
||
|
|
|
||
|
|
plt.figure(figsize=(10, 6))
|
||
|
|
plt.plot(accuracies, label='Acurácia por Imagem')
|
||
|
|
plt.xlabel('Número da Imagem')
|
||
|
|
plt.ylabel('Acurácia')
|
||
|
|
plt.title('Acurácia por Imagem no Conjunto de Teste')
|
||
|
|
plt.legend()
|
||
|
|
plt.show()
|
||
|
|
|
||
|
|
print("Validação concluída.")
|