Implementação e Análise do Algoritmo ResNet50V2 em PyTorch

Este artigo explora a implementação do modelo ResNet50V2 usando a biblioteca PyTorch, detalhando sua construção, treinamento e as principais diferenças em relação ao ResNetV1. Serão abordados desde a configuração inicial do ambiente e o pré-processamento de dados até a arquitetura da rede neural e o processo de treinamento.

1. Configuração do Ambiente e Preparação de Dados

Inicialmente, configuramos o ambiente de desenvolvimento e preparamos os dados para o treinamento do modelo.

1.1. Configuração da GPU

Verificamos a disponibilidade de uma GPU para acelerar o treinamento do modelo, utilizando-a se detectada, caso contrário, o processamento será feito na CPU.

import torch
import torch.nn as nn
import torchvision
from torchvision import transforms, datasets
from torch.utils.data import DataLoader, random_split

import os
import pathlib
import matplotlib.pyplot as plt
from PIL import Image

# Configura o dispositivo para GPU se disponível, caso contrário CPU
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(f"Dispositivo de processamento: {device}")

1.2. Carregamento e Pré-processamento dos Dados

Os dados de imagem são carregados de um diretório específico. Realizamos transformações essenciais como redimensionamento, conversão para tensor e normalização, seguindo os padrões de pré-processamento para modelos de visão computacional pré-treinados.

# Define o diretório dos dados
caminho_dados = pathlib.Path("F:/365data/J3/")

# Lista as classes presentes no diretório
nomes_classes = [str(pasta).split('\\')[3] for pasta in caminho_dados.glob('*')]
print(f"Classes identificadas: {nomes_classes}")

# Define as transformações para treinamento e validação/teste
transformacoes_treino = transforms.Compose([
    transforms.Resize([224, 224]), # Redimensiona todas as imagens para 224x224
    transforms.ToTensor(), # Converte as imagens para tensores PyTorch
    transforms.Normalize( # Normaliza as imagens com base nas estatísticas do ImageNet
        mean=[0.485, 0.456, 0.406],
        std=[0.229, 0.224, 0.225]
    )
])

transformacoes_teste = transforms.Compose([
    transforms.Resize([224, 224]),
    transforms.ToTensor(),
    transforms.Normalize(
        mean=[0.485, 0.456, 0.406],
        std=[0.229, 0.224, 0.225]
    )
])

# Carrega o conjunto de dados completo
conjunto_total = datasets.ImageFolder(str(caminho_dados), transform=transformacoes_treino)
print(f"Total de imagens carregadas: {len(conjunto_total)}")

# Divide o conjunto de dados em treinamento e teste (80/20)
tamanho_treino = int(0.8 * len(conjunto_total))
tamanho_teste = len(conjunto_total) - tamanho_treino
conjunto_treino, conjunto_teste = random_split(conjunto_total, [tamanho_treino, tamanho_teste])

print(f"Imagens para treino: {len(conjunto_treino)}")
print(f"Imagens para teste: {len(conjunto_teste)}")

1.3. Criação dos DataLoaders

Os DataLoaders são criados para gerenciar o carregamento dos dados em lotes (batches), facilitando o processo de treinamento e avaliação.

tamanho_lote = 32

dataloader_treino = DataLoader(conjunto_treino,
                                       batch_size=tamanho_lote,
                                       shuffle=True,
                                       num_workers=os.cpu_count()) # Usar todos os cores disponíveis

dataloader_teste = DataLoader(conjunto_teste,
                                      batch_size=tamanho_lote,
                                      shuffle=False, # Não é necessário embaralhar o conjunto de teste
                                      num_workers=os.cpu_count())

# Verifica a forma dos tensores de entrada e rótulos
for imagens, rotulos in dataloader_teste:
    print(f'Forma das imagens: {imagens.shape}')
    print(f'Forma dos rótulos: {rotulos.shape}, Tipo: {rotulos.dtype}')
    break

1.4. Visualização de Dados

Para inspecionar o conjunto de dados, exibimos algumas imagens de exemplo.

# Exemplo de visualização de algumas imagens
diretorio_exemplo_imagens = 'F:/365data/J1/bird_photos/Bananaquit/' # Assumindo um diretório similar para visualização

if os.path.exists(diretorio_exemplo_imagens):
    arquivos_imagens = [f for f in os.listdir(diretorio_exemplo_imagens) if f.endswith(('.jpg', '.jpeg', '.png'))]
    fig, axes = plt.subplots(1, 5, figsize=(10, 3)) # Exibe 5 imagens
    for i, ax in enumerate(axes.flat):
        if i < len(arquivos_imagens):
            caminho_imagem = os.path.join(diretorio_exemplo_imagens, arquivos_imagens[i])
            img = Image.open(caminho_imagem)
            ax.imshow(img)
            ax.axis('off')
        else:
            ax.remove() # Remove eixos vazios se houver menos de 5 imagens
    plt.tight_layout()
    plt.show()
else:
    print(f"Diretório de exemplo de imagens não encontrado: {diretorio_exemplo_imagens}")

2. Construção do Modelo ResNet50V2

O ResNet50V2 é uma variação da arquitetura ResNet que incorpora pré-ativação (Batch Normalization e ReLU antes da convolução) e, conforme a descrição original, um tipo de atalho via max-pooling em blocos específicos. O modelo usa blocos residuais tipo "bottleneck".

2.1. Definição dos Blocos Residuais (Bottleneck)

Os blocos residuais são a base da arquitetura ResNet. Para ResNetV2, a sequência BN-ReLU é aplicada antes de cada camada convolucional no caminho principal, e o atalho opera sobre a entrada original do bloco, sem uma ativação final pós-soma.

Aqui, definimos três tipos de blocos bottleneck para corresponder à estrutura do modelo original: um para projeção (mudança de dimensão e canais via convolução no atalho), um para identidade (canais e dimensão iguais, atalho de identidade), e um especial com atalho de max-pooling para downsampling espacial, conforme mencioando na descrição do ResNet50V2.

class BlocoBottleneckV2(nn.Module):
    expansao = 4 # Fator de expansão para os canais de saída do bloco

    def __init__(self, canais_entrada, canais_intermediarios, stride=1, tipo_atalho='projecao'):
        super().__init__()

        self.tipo_atalho = tipo_atalho
        self.stride = stride
        canais_saida_bloco = canais_intermediarios * self.expansao

        # Caminho principal com pré-ativação (BN-ReLU antes de cada Conv)
        self.sequencia_conv = nn.Sequential(
            nn.BatchNorm2d(canais_entrada),
            nn.ReLU(inplace=True),
            nn.Conv2d(canais_entrada, canais_intermediarios, kernel_size=1, bias=False),

            nn.BatchNorm2d(canais_intermediarios),
            nn.ReLU(inplace=True),
            nn.Conv2d(canais_intermediarios, canais_intermediarios, kernel_size=3, stride=stride, padding=1, bias=False),

            nn.BatchNorm2d(canais_intermediarios),
            nn.ReLU(inplace=True),
            nn.Conv2d(canais_intermediarios, canais_saida_bloco, kernel_size=1, bias=False)
        )

        # Configuração do atalho
        if tipo_atalho == 'projecao':
            # Atalho convolucional para corresponder canais e possivelmente downsample espacial
            self.atalho = nn.Conv2d(canais_entrada, canais_saida_bloco, kernel_size=1, stride=stride, bias=False)
        elif tipo_atalho == 'identidade':
            # Atalho de identidade, assumindo que canais_entrada já é igual a canais_saida_bloco e stride=1
            self.atalho = nn.Identity()
        elif tipo_atalho == 'maxpool_especial':
            # Atalho de MaxPool, como descrito para o ResNet50V2 original
            # Assumimos que canais_entrada == canais_saida_bloco e que o MaxPool serve para downsampling espacial
            self.atalho = nn.MaxPool2d(kernel_size=1, stride=stride)
        else:
            raise ValueError("Tipo de atalho desconhecido.")

    def forward(self, x):
        caminho_residual = x
        caminho_principal = self.sequencia_conv(x)

        # Aplica o atalho ao caminho residual
        caminho_residual_processado = self.atalho(caminho_residual)

        return caminho_principal + caminho_residual_processado # ResNetV2 não tem ReLU final no bloco

2.2. Definição do Modelo ResNet50V2 Completo

O modelo ResNet50V2 é construído empilhando esses blocos residuais em estágios. A camada inicial (stem) e as camadas finais de classificação são também incluídas.

class ResNet50V2(nn.Module):
    def __init__(self, num_classes=len(nomes_classes)):
        super().__init__()

        # Camada inicial (Stem)
        self.camada_inicial = nn.Sequential(
            nn.ZeroPad2d(3),
            nn.Conv2d(3, 64, kernel_size=7, stride=2, bias=False),
            nn.BatchNorm2d(64),
            nn.ReLU(inplace=True),
            nn.ZeroPad2d(1),
            nn.MaxPool2d(kernel_size=3, stride=2)
        )

        # Definição dos estágios do ResNet50V2, utilizando os blocos customizados
        # Estágio 1 (64 canais de entrada -> 256 canais de saída)
        # O ResNet50V2 original usava a sequência de 3 blocos como `block1`, `block2`, `block3`.
        # `block1`: projeção (in=64, inter=64 -> out=256), stride=1
        self.estagio1_bloco1 = BlocoBottleneckV2(64, 64, stride=1, tipo_atalho='projecao')
        # `block2`: identidade (in=256, inter=64 -> out=256), stride=1
        self.estagio1_bloco2 = BlocoBottleneckV2(256, 64, stride=1, tipo_atalho='identidade')
        # `block3`: maxpool_especial (in=256, inter=64 -> out=256), stride=2 (downsampling)
        self.estagio1_bloco3 = BlocoBottleneckV2(256, 64, stride=2, tipo_atalho='maxpool_especial') # Downsampling aqui

        # Estágio 2 (256 canais de entrada -> 512 canais de saída)
        # `block4`: projeção (in=256, inter=128 -> out=512), stride=1 (mas prepara para downsampling na próxima camada)
        self.estagio2_bloco1 = BlocoBottleneckV2(256, 128, stride=1, tipo_atalho='projecao')
        self.estagio2_bloco2 = BlocoBottleneckV2(512, 128, stride=1, tipo_atalho='identidade')
        self.estagio2_bloco3 = BlocoBottleneckV2(512, 128, stride=1, tipo_atalho='identidade')
        # `block7`: maxpool_especial (in=512, inter=128 -> out=512), stride=2
        self.estagio2_bloco4 = BlocoBottleneckV2(512, 128, stride=2, tipo_atalho='maxpool_especial')

        # Estágio 3 (512 canais de entrada -> 1024 canais de saída)
        # `block8`: projeção (in=512, inter=256 -> out=1024), stride=1
        self.estagio3_bloco1 = BlocoBottleneckV2(512, 256, stride=1, tipo_atalho='projecao')
        self.estagio3_bloco2 = BlocoBottleneckV2(1024, 256, stride=1, tipo_atalho='identidade')
        self.estagio3_bloco3 = BlocoBottleneckV2(1024, 256, stride=1, tipo_atalho='identidade')
        self.estagio3_bloco4 = BlocoBottleneckV2(1024, 256, stride=1, tipo_atalho='identidade')
        self.estagio3_bloco5 = BlocoBottleneckV2(1024, 256, stride=1, tipo_atalho='identidade')
        # `block13`: maxpool_especial (in=1024, inter=256 -> out=1024), stride=2
        self.estagio3_bloco6 = BlocoBottleneckV2(1024, 256, stride=2, tipo_atalho='maxpool_especial')

        # Estágio 4 (1024 canais de entrada -> 2048 canais de saída)
        # `block14`: projeção (in=1024, inter=512 -> out=2048), stride=1
        self.estagio4_bloco1 = BlocoBottleneckV2(1024, 512, stride=1, tipo_atalho='projecao')
        self.estagio4_bloco2 = BlocoBottleneckV2(2048, 512, stride=1, tipo_atalho='identidade')
        self.estagio4_bloco3 = BlocoBottleneckV2(2048, 512, stride=1, tipo_atalho='identidade')

        # Camadas finais de classificação
        self.final_preativacao_pool = nn.Sequential(
            nn.BatchNorm2d(2048),
            nn.ReLU(inplace=True),
            nn.AdaptiveAvgPool2d((1, 1)) # Global Average Pooling
        )
        self.classificador = nn.Linear(2048, num_classes)

        # Inicialização dos pesos
        self._inicializar_pesos()

    def _inicializar_pesos(self):
        for m in self.modules():
            if isinstance(m, nn.Conv2d):
                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
            elif isinstance(m, nn.BatchNorm2d):
                nn.init.constant_(m.weight, 1)
                nn.init.constant_(m.bias, 0)
            elif isinstance(m, nn.Linear):
                nn.init.constant_(m.bias, 0)

    def forward(self, x):
        x = self.camada_inicial(x)

        x = self.estagio1_bloco1(x)
        x = self.estagio1_bloco2(x)
        x = self.estagio1_bloco3(x)

        x = self.estagio2_bloco1(x)
        x = self.estagio2_bloco2(x)
        x = self.estagio2_bloco3(x)
        x = self.estagio2_bloco4(x)

        x = self.estagio3_bloco1(x)
        x = self.estagio3_bloco2(x)
        x = self.estagio3_bloco3(x)
        x = self.estagio3_bloco4(x)
        x = self.estagio3_bloco5(x)
        x = self.estagio3_bloco6(x)

        x = self.estagio4_bloco1(x)
        x = self.estagio4_bloco2(x)
        x = self.estagio4_bloco3(x)

        x = self.final_preativacao_pool(x)
        x = torch.flatten(x, 1) # Equivalente a x.view(x.size(0), -1)
        x = self.classificador(x)

        return x

# Instancia o modelo e o move para o dispositivo selecionado (GPU/CPU)
modelo_resnet50v2 = ResNet50V2().to(device)
print(modelo_resnet50v2)

Diferenças entre ResNetV2 e ResNetV1: A principle inovação do ResNetV2, conforme implementado aqui, é a pré-ativação, onde as camadas de Batch Normalization (BN) e ReLU são aplicadas antes da camada convolucional. Isso contrasta com o ResNetV1, que geralmente aplica BN após a convolução e ReLU após a soma do atalho. Além disso, esta implementação inclui um tipo de bloco residual com um atalho que utiliza MaxPool, uma característica distinta que difere dos atalhos convolucionais ou de identidade padrão encontrados no ResNetV1.

2.3. Resumo dos Parâmetros do Modelo

Utilizamos a biblioteca torchsummary para exibir um resumo detalhado do modelo, incluindo o número total de parâmetros treináveis.

import torchsummary as summary
# Para usar torchsummary, é necessário garantir que o modelo e o input estão no mesmo device.
summary.summary(modelo_resnet50v2, (3, 224, 224), device=str(device))

3. Treinamento do Modelo

Definimos as funções de trienamento e teste, configuramos o otimizador e a função de perda, e executamos o ciclo de treinamento.

3.1. Função de Treinamento

Esta função itera sobre o DataLoader de treinamento, calcula a perda, realiza a retropropagação e atualiza os pesos do modelo.

def treinar_epoca(dataloader, modelo, otimizador, funcao_perda):
    tamanho_dataset = len(dataloader.dataset)
    num_lotes = len(dataloader)
    
    perda_total_treino, acuracia_total_treino = 0, 0
    
    modelo.train() # Coloca o modelo em modo de treinamento
    for X_lote, y_lote in dataloader:
        X_lote, y_lote = X_lote.to(device), y_lote.to(device)
        
        # Computa a previsão e a perda
        predicoes = modelo(X_lote)
        perda = funcao_perda(predicoes, y_lote)
        
        # Retropropagação
        otimizador.zero_grad() # Zera os gradientes acumulados
        perda.backward() # Computa os gradientes
        otimizador.step() # Atualiza os pesos
        
        perda_total_treino += perda.item()
        acuracia_total_treino += (predicoes.argmax(1) == y_lote).type(torch.float).sum().item()
        
    perda_media_treino = perda_total_treino / num_lotes
    acuracia_media_treino = acuracia_total_treino / tamanho_dataset
    
    return acuracia_media_treino, perda_media_treino

3.2. Função de Teste/Avaliação

Esta função avalia o desempenho do modelo em um conjunto de dados de teste, sem realizar atualizações de pesos.

def testar_epoca(dataloader, modelo, funcao_perda):
    tamanho_dataset = len(dataloader.dataset)
    num_lotes = len(dataloader)
    perda_total_teste, acuracia_total_teste = 0, 0
    
    modelo.eval() # Coloca o modelo em modo de avaliação
    with torch.no_grad(): # Desativa o cálculo de gradientes para economizar memória e computação
        for X_lote, y_lote in dataloader:
            X_lote, y_lote = X_lote.to(device), y_lote.to(device)
            
            predicoes = modelo(X_lote)
            perda = funcao_perda(predicoes, y_lote)
            
            perda_total_teste += perda.item()
            acuracia_total_teste += (predicoes.argmax(1) == y_lote).type(torch.float).sum().item()
            
    perda_media_teste = perda_total_teste / num_lotes
    acuracia_media_teste = acuracia_total_teste / tamanho_dataset
    
    return acuracia_media_teste, perda_media_teste

3.3. Configuração do Otimizador e Função de Perda

Utilizamos a função de perda CrossEntropyLoss, adequada para problemas de classificação multi-classe, e o otimizador SGD (Stochastic Gradient Descent).

funcao_perda = nn.CrossEntropyLoss()
taxa_aprendizagem = 1e-3
otimizador = torch.optim.SGD(modelo_resnet50v2.parameters(), lr=taxa_aprendizagem)

3.4. Ciclo de Treinamento Completo

O modelo é treinado por um número definido de épocas, registrando as métricas de desempenho e salvando o melhor modelo.

import copy
import warnings
warnings.filterwarnings("ignore") # Ignora warnings para facilitar a leitura da saída

num_epocas = 20

historico_perda_treino = []
historico_acuracia_treino = []
historico_perda_teste = []
historico_acuracia_teste = []
melhor_acuracia = 0.0
melhor_modelo_estado = None

print("Iniciando treinamento...")
for epoca in range(num_epocas):
    acuracia_treino_epoca, perda_treino_epoca = treinar_epoca(dataloader_treino, modelo_resnet50v2, otimizador, funcao_perda)
    acuracia_teste_epoca, perda_teste_epoca = testar_epoca(dataloader_teste, modelo_resnet50v2, funcao_perda)

    # Salva o melhor modelo com base na acurácia de teste
    if acuracia_teste_epoca > melhor_acuracia:
        melhor_acuracia = acuracia_teste_epoca
        melhor_modelo_estado = copy.deepcopy(modelo_resnet50v2.state_dict())
        
    historico_acuracia_treino.append(acuracia_treino_epoca)
    historico_perda_treino.append(perda_treino_epoca)
    historico_acuracia_teste.append(acuracia_teste_epoca)
    historico_perda_teste.append(perda_teste_epoca)
    
    lr_atual = otimizador.param_groups[0]['lr'] # Obtém a taxa de aprendizagem atual
    
    template = ('Época: {:2d}, Acurácia Treino: {:.1f}%, Perda Treino: {:.3f}, Acurácia Teste: {:.1f}%, Perda Teste: {:.3f}, LR: {:.2E}')
    print(template.format(epoca + 1, acuracia_treino_epoca * 100, perda_treino_epoca, 
                          acuracia_teste_epoca * 100, perda_teste_epoca, lr_atual))

# Salva o estado do melhor modelo treinado
caminho_salvamento_modelo = 'F:/365data/J2_melhor_modelo.pth' # Certifique-se de que o diretório existe
if melhor_modelo_estado:
    torch.save(melhor_modelo_estado, caminho_salvamento_modelo)
    print(f"Melhor modelo salvo em: {caminho_salvamento_modelo}")
else:
    print("Nenhum modelo foi salvo (melhor acurácia de teste não foi alcançada).")

print('Treinamento concluído!')

Exemplo de Saída do Treinamento:

Época:  1, Acurácia Treino: 41.6%, Perda Treino: 1.387, Acurácia Teste: 28.3%, Perda Teste: 144.182, LR: 1.00E-03
Época:  2, Acurácia Treino: 60.8%, Perda Treino: 0.884, Acurácia Teste: 25.7%, Perda Teste: 8.248, LR: 1.00E-03
Época:  3, Acurácia Treino: 78.3%, Perda Treino: 0.615, Acurácia Teste: 44.2%, Perda Teste: 2.175, LR: 1.00E-03
...
Época: 17, Acurácia Treino: 97.3%, Perda Treino: 0.166, Acurácia Teste: 85.0%, Perda Teste: 0.396, LR: 1.00E-03
...
Época: 20, Acurácia Treino: 97.8%, Perda Treino: 0.087, Acurácia Teste: 85.0%, Perda Teste: 0.426, LR: 1.00E-03

3.5. Visualização do Histórico de Treinamento

Os gráficos de acurácia e perda de treinamento e teste são gerados para visualizar o desempenho do modelo ao longo das épocas.

# Configurações para exibir gráficos
plt.rcParams['font.sans-serif'] = ['SimHei'] # Para exibir caracteres chineses, se necessário
plt.rcParams['axes.unicode_minus'] = False # Para exibir sinais negativos
plt.rcParams['figure.dpi'] = 100 # Resolução da figura

epocas_range = range(num_epocas)

plt.figure(figsize=(12, 4))
plt.subplot(1, 2, 1)
plt.plot(epocas_range, historico_acuracia_treino, label='Acurácia de Treinamento')
plt.plot(epocas_range, historico_acuracia_teste, label='Acurácia de Teste')
plt.legend(loc='lower right')
plt.title('Acurácia de Treinamento e Teste')

plt.subplot(1, 2, 2)
plt.plot(epocas_range, historico_perda_treino, label='Perda de Treinamento')
plt.plot(epocas_range, historico_perda_teste, label='Perda de Teste')
plt.legend(loc='upper right')
plt.title('Perda de Treinamento e Teste')
plt.show()

Tags: Pytorch ResNet50V2 DeepLearning ComputerVision ConvolutionalNeuralNetworks

Publicado em 9-1 09:42