import torch from PIL import Image import torchvision.transforms as T import matplotlib.pyplot as plt import numpy as np # Supondo que você tenha a estrutura do repositório e o módulo `network` conforme descrito no README from network.modeling import deeplabv3plus_resnet50 as deeplabv3_model # Configurações Iniciais NUM_CLASSES = 4 # Pascal VOC possui 3 classes + 1 para o fundo OUTPUT_STRIDE = 16 # Valor comum para DeepLab MODEL_PATH = 'backup/ruasModel_final.pth' # Caminho para o modelo pré-treinado # Função para carregar o modelo def load_model(model_path): model = deeplabv3_model(num_classes=NUM_CLASSES, output_stride=OUTPUT_STRIDE) model.load_state_dict(torch.load(model_path), strict=False) model.eval() # Modo de avaliação return model # Carregar o modelo model = load_model(MODEL_PATH) # Função para processar uma imagem e realizar a segmentação def segment_image(image_path, model): image = Image.open(image_path).convert("RGB") transform = T.Compose([ T.Resize(520), T.ToTensor(), T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) input_tensor = transform(image).unsqueeze(0) with torch.no_grad(): output = model(input_tensor) output_predictions = output.max(1)[1].squeeze().detach().cpu().numpy() return output_predictions def read_labelmap(path): label_colors = [] class_names = [] with open(path, 'r') as file: for line in file.readlines(): # Ignora linhas comentadas if line.startswith('#'): continue parts = line.strip().split(':') if len(parts) >= 2: label = parts[0].strip() color = tuple(map(int, parts[1].split(','))) class_names.append(label) label_colors.append(color) return np.array(label_colors), class_names def decode_segmap(image, labelmap_path="dataset/labelmap.txt"): label_colors, class_names = read_labelmap(labelmap_path) nc = len(label_colors) detected_classes = set(np.unique(image)) # Conjunto das classes detectadas na imagem r = np.zeros_like(image).astype(np.uint8) g = np.zeros_like(image).astype(np.uint8) b = np.zeros_like(image).astype(np.uint8) for l in detected_classes: # Usar apenas as classes detectadas if l < nc: # Verifica se o índice está dentro do intervalo das cores definidas idx = image == l r[idx] = label_colors[l, 0] g[idx] = label_colors[l, 1] b[idx] = label_colors[l, 2] rgb = np.stack([r, g, b], axis=2) return rgb, class_names, label_colors, detected_classes def display_segmentation(input_image_path, output_predictions): original_image = Image.open(input_image_path) original_image = np.array(original_image) seg_image, class_names, label_colors, detected_classes = decode_segmap(output_predictions) seg_image_resized = Image.fromarray(seg_image).resize((original_image.shape[1], original_image.shape[0]), resample=Image.NEAREST) seg_image_resized = np.array(seg_image_resized) overlayed_img = original_image * 0.6 + seg_image_resized * 0.4 overlayed_img = overlayed_img.astype(np.uint8) plt.figure(figsize=(20, 10)) # Imagem Original plt.subplot(1, 3, 1) plt.imshow(original_image) plt.title('Imagem Original') plt.axis('off') # Representação da Segmentação plt.subplot(1, 3, 2) plt.imshow(seg_image_resized) plt.title('Representação da Segmentação') plt.axis('off') # Imagem com Overlay plt.subplot(1, 3, 3) plt.imshow(overlayed_img) plt.title('Imagem com Overlay') plt.axis('off') import matplotlib.patches as mpatches patches = [mpatches.Patch(color=np.array(label_colors[i])/255.0, label=class_names[i]) for i in detected_classes] plt.figlegend(handles=patches, loc='lower center', ncol=5, labelspacing=0.) plt.subplots_adjust(bottom=0.25) # Ajuste conforme necessário para acomodar a legenda plt.show() image_path = 'dataset/val_images/7.jpg' # Exemplo de como usar a função `segment_image` output_predictions = segment_image(image_path, model) # Chamada da função display_segmentation(image_path, output_predictions)