O treinamento de redes neurais profundas enfrenta desafios significativos relacionados à propagação de gradiantes. Durante o processo de backpropagation, os gradientes podem sofrer variações drásticas de escala, levando aos problemas conhecidos como "explosão de gradiente" ou "desaparecimento de gradiente". Esses fenômenos resultam em convergência lenta ou instabilidade no treinamento. A Normalização em Lote (Batch Normalization ou BN) surge como uma técnica eficaz para mitigar esses problemas, restringindo a distribuição das ativações intermediárias a um intervalo controlado.
Mecanismo da Normalização em Lote
A BN opera normalizando as saídas de uma camada anterior para cada mini-lote (mini-batch). O processo consiste em calcular a média e a variância dos dados no lote atual e padronizar os valores. Matematicamente, para um vetor de entrada \(X\), a normalização é definida como:
\[ \hat{X}_i = \frac{X_i - \mu_B}{\sqrt{\sigma_B^2 + \epsilon}} \]
Onde \(\mu_B\) representa a média do lote e \(\sigma_B^2\) a variância. O termo \(\epsilon\) é uma constante pequena adicionada para evitar divisão por zero. Essa operação força os dados a terem média zero e variância unitária, o que mantém os gradientes em uma faixa saudável para funções de ativação como a sigmoide, onde a derivada é mais expressiva próximo à origem.
No entanto, simplesmente normalizar os dados pode reduzir a capacidade de representação da rede. A força das redes neurais reside na não-linearidade. Para evitar que as ativações fiquem restritas a uma região excessivamente linear, a BN introduz dois parâmetros aprendíveis, \(\gamma\) (escala) e \(\beta\) (deslocamento), permitindo que o modelo recupere a representação original se necessário:
\[ y_i = \gamma \hat{X}_i + \beta \]
Esses parâmetros são otimizados via backpropagation, perimtindo que a rede ajuste a distribuição normalizada de forma ideal para a tarefa.
Durante a inferência, o modelo processa amostras individuais, tornando impossível o cálculo da média e variância do lote. Para contornar isso, a BN utiliza estatísticas acumuladas (média móvel) de todos os lotes vistos durante o treinamento para normalizar a entrada única.
Implementação em Keras
O framework Keras oferece uma implementação pronta e otimizada para BN. O código abaixo demonstra sua aplicação em um tensor de entrada, tipicamente utilizado após camadas convolucionais:
import tensorflow as tf
from tensorflow.keras import layers, Model, Input
# Definição da entrada (ex: imagem 320x320 com 3 canais)
entrada = Input(shape=(320, 320, 3))
# Aplicação da Normalização em Lote
# axis=-1 indica normalização ao longo do eixo dos canais
camada_bn = layers.BatchNormalization(
axis=-1,
momentum=0.99,
epsilon=0.001,
center=True, # Aplica o parâmetro beta
scale=True # Aplica o parâmetro gamma
)(entrada)
modelo = Model(inputs=entrada, outputs=camada_bn)
modelo.summary()
Ao analisar o resumo do modelo, notam-se parâmetros treináveis (\(\gamma\) e \(\beta\)) e não treináveis (a média móvel e a variância). É crucial entender o comportamento do argumento trainable. Se definido como False, os parâmetros \(\gamma\) e \(\beta\) são congelados. Se True, eles são atualizados via gradiente, enquanto as estatísticas móveis continuam sendo atualizadas internamente pelo momento, independentemente do estado de treinamento do gradiente.
Implementação em PyTorch
No PyTorch, a lógica é similar, mas a API difere levemente. Abaixo, apresentamos um exemplo utilizando a classe nn.BatchNorm2d, adequada para entradas convolucionais 4D (Batch, Canais, Altura, Largura):
import torch
import torch.nn as nn
# Instanciação da camada
# num_features corresponde ao número de canais (C)
bn_pytorch = nn.BatchNorm2d(
num_features=3,
eps=1e-05,
momentum=0.1, # Valor padrão, usado para média móvel
affine=True, # Habilita aprendizado de gamma e beta
track_running_stats=True
)
# Simulação de entrada
dados_entrada = torch.randn(16, 3, 64, 64)
saida = bn_pytorch(dados_entrada)
# Verificação dos buffers internos (média e variância acumuladas)
for nome, param in bn_pytorch.named_buffers():
print(f"{nome}: {param}")
Um ponto de atenção no PyTorch é o argumento track_running_stats. Quando True, a camada atualiza suas estatísticas internas durante o treinamento. Durante a avaliação (model.eval()), a camada utiliza essas estatísticas fixas. Se desativado, a normalização na inferência dependerá das estatísticas do batch de entrada, o que geralmente não é desejado.