Contexto e Limitações de Hardware
Durante a análise de precisão com ferramentas de depuração de quantização, é comum identificar operações lineares críticas. Suponha que os diagnósticos apontem que os pesos de uma camada específica exigem precisão de 16 bits (int16) devido à sua alta sensibilidade, apresentando uma distribuição estatística com variância elevada e valores extremos. O desafio técnico surge quando a arquitetura de hardware alvo, como a BPU dos processadores da família Jornada 6 (J6E/M), possui uma restrição intrínseca: não há suporte nativo para operações Linear onde tanto a entrada (input) quanto os pesos (weight) sejam quantizados simultaneamente em int16.
Estratégia de Decomposição Matemática
Para contornar essa limitação de hardware sem a necessidade de retreinar o modelo original em ponto flutuante, a camada linear sensível pode ser decomposta em operações primitivas suportadas. Matematicamente, uma camada linear é equivalente a uma multiplicação com difusão (broadcast mul) seguida de uma redução por soma (sum).
Diferenças no Paradigma de Quantização:
- Em uma operação
Linearpadrão, os pesos são processados utilizando quentização por canal (per-chennel). - Ao tratar os pesos como tensores de entrada para uma operação de multiplicação, a quantização passa a ser aplicada por tensor (per-tensor).
Na prática, migrar a representação dos pesos de int8 per-channel para int16 per-tensor costuma promover uma otimização positiva na acurácia geral. Este ajuste estrutural deve ser aplicado no grafo computacional após a fase de treinamento em ponto flutuante, precedendo as etapas de calibração e QAT (Quantization-Aware Training).
Armadilhas na Calibração MSE e Mitigação
Ao adotar a abordagem de multiplicação e soma, é necessário monitorar a distribuição dos dados intermediários. Se a grande maioria dos valores resultantes da multiplicação estiver concentrada próxima a zero, a calibração baseada em Erro Quadrático Médio (MSE) torna-se vulnerável a valores atípicos (outliers). Isso resulta em uma escala (scale) de saída excessivamente inflada, forçando o arredondamento de pequenos valores para zero absoluto e causando um desvio drástico na operação de soma subsequente.
Mitigação: Este cenário é particularmente crítico quando a multiplicação precede funções de ativação como a Sigmoid. A solução recomendada é fixar a escala de saída da operação de multiplicação (por exemplo, definindo um valor fixo como 7/32767), garantindo que a faixa dinâmica necessária para a Sigmoid seja preservada sem distorções severas.
Implementação Prática e Refatoração
O exemplo abaixo demonstra a substituição estrutural da camada linear sensível. A lógica foi refatorada para utilizar operações de expansão de dimensões (unsqueeze) dinâmicas, eliminando dependências de formatos de lote rígidos e encapsulando os pesos originais como parâmetros não treináveis.
import torch
import torch.nn as nn
from horizon_plugin_pytorch import set_march, March
from horizon_plugin_pytorch.quantization import prepare, set_fake_quantize, FakeQuantState, QuantStub
from horizon_plugin_pytorch.quantization.hbdk4 import export
from horizon_plugin_pytorch.quantization.qconfig_template import calibration_8bit_weight_16bit_act_qconfig_setter
from torch.quantization import DeQuantStub
class NetworkWithBypassedLinear(nn.Module):
def __init__(self, target_weights, target_biases):
super().__init__()
# Blocos iniciais do modelo
self.dense_1 = nn.Linear(256, 256)
self.norm_layer = nn.LayerNorm(256)
self.activation_fn = nn.ReLU()
# Injeção dos pesos e vieses da camada sensível original
self.register_parameter('sensitive_w', nn.Parameter(target_weights, requires_grad=False))
self.register_parameter('sensitive_b', nn.Parameter(target_biases, requires_grad=False))
# Bloco subsequente
self.dense_3 = nn.Linear(60, 60)
# Stubs de controle de quantização
self.input_quantizer = QuantStub()
self.weight_quantizer = QuantStub()
self.bias_quantizer = QuantStub()
self.output_dequantizer = DeQuantStub()
def forward(self, inputs):
x = self.input_quantizer(inputs)
q_weights = self.weight_quantizer(self.sensitive_w)
q_biases = self.bias_quantizer(self.sensitive_b)
# Processamento inicial
x = self.dense_1(x)
x = self.norm_layer(x)
x = self.activation_fn(x)
# Substituição da Camada Linear Sensível via Broadcast Mul + Sum
# Expansão dinâmica para compatibilidade de broadcast
# x shape: [Batch, Seq, Features_In] -> [Batch, Seq, 1, Features_In]
x_expanded = x.unsqueeze(2)
# weights shape: [Features_Out, Features_In] -> [1, 1, Features_Out, Features_In]
w_expanded = q_weights.unsqueeze(0).unsqueeze(1)
# Multiplicação elemento a elemento com difusão
broadcast_product = x_expanded * w_expanded
# Redução por soma ao longo da dimensão de entrada original
weighted_sum = broadcast_product.sum(dim=-1)
# Aplicação do termo de bias
x = weighted_sum + q_biases
# Processamento final
x = self.dense_3(x)
return self.output_dequantizer(x)
# Configuração do ambiente de hardware alvo
set_march(March.NASH_M)
# Simulação de carregamento de tensores de validação
sample_inputs = torch.randn(2, 100, 256)
dummy_weights = torch.randn(60, 256)
dummy_biases = torch.randn(60)
model_instance = NetworkWithBypassedLinear(dummy_weights, dummy_biases)
# Preparação do pipeline de calibração
calib_model = prepare(
model_instance.eval(),
sample_inputs,
qconfig_setter=calibration_8bit_weight_16bit_act_qconfig_setter
)
# Execução da fase de calibração estatística
set_fake_quantize(calib_model, FakeQuantState.CALIBRATION)
with torch.no_grad():
calib_model(sample_inputs)
# Execução da fase de validação quantizada
set_fake_quantize(calib_model, FakeQuantState.VALIDATION)
with torch.no_grad():
calibrated_output = calib_model(sample_inputs)
# Exportação do modelo para o backend BPU
quantized_backend_model = export(calib_model, sample_inputs)
# final_compiled_model = convert(quantized_backend_model, March.NASH_M)
Análise de Consistência dos Tensores
A decomposição matemática preserva a fidelidade numérica do modelo original. Ao comaprar a saída em ponto flutuante (Float) com a saída após a calibração (Calib) utilizando a topologia baseada em multiplicação e soma, a divergência mínima demonstra a eficácia do método para contornar a restrição de int16 simultâneo na BPU.
Referência de Saída em Ponto Flutuante:
tensor([[[-0.3112, 0.1458, -0.5111, ..., -0.0621, -0.2213, -0.0312],
[-0.2019, -0.0115, -0.3301, ..., 0.3104, -0.0909, -0.0587],
[-0.3076, 0.1512, -0.2613, ..., 0.2401, -0.3512, 0.0154],
...,
[-0.4010, -0.0411, -0.1726, ..., -0.0088, -0.4191, 0.0512],
[-0.1109, 0.2501, -0.1916, ..., 0.0812, -0.3501, 0.0251],
[-0.2214, -0.1121, -0.0615, ..., 0.3514, -0.1721, 0.2318]]],
grad_fn=<ViewBackward0>)
Saída Calibrada com Substituição (Broadcast Mul + Sum Int16):
tensor([[[-0.3138, 0.1480, -0.5129, ..., -0.0641, -0.2231, -0.0299],
[-0.2035, -0.0091, -0.3310, ..., 0.3125, -0.0924, -0.0591],
[-0.3092, 0.1531, -0.2637, ..., 0.2418, -0.3539, 0.0172],
...,
[-0.4028, -0.0423, -0.1742, ..., -0.0095, -0.4217, 0.0524],
[-0.1128, 0.2522, -0.1930, ..., 0.0820, -0.3527, 0.0260],
[-0.2232, -0.1141, -0.0627, ..., 0.3532, -0.1737, 0.2335]]],
grad_fn=<ViewBackward0>)