Implementação de Quantização Int16 Dupla em Camadas Lineares Sensíveis no Horizon J6

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 Linear padrã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>)

Tags: Pytorch HorizonJ6 Quantização Int16 BPU

Publicado em 9-5 04:11