Flujo de despliegue de redes neuronales con TVM: del algoritmo al runtime

Flujo de trabajo de TVM

Para desplegar un modelo de aprendizaje profundo sobre un nuevo acelerador, TVM sigue una cadena de compilación que va desde la representación del grafo hasta el código máquina. Los pasos principales son:

  • Importación del modelo: se lee el pesos y la arquitectura desde PyTorch, TensorFlow, ONNX u otros frameworks.
  • Conversión a Relay: Relay es la representación intermedia de alto nivel. La versión más reciente es Relax.
  • Conversión a TE: tras las optimizaciones de grafo, los subgrafos se bajan a Tensor Expression (TE).
  • Autoajuste: AutoTVM o AutoScheduler buscan el plan de ejecución más eficiente para la plataforma objetivo.
  • Conversión a TIR: Tensor Intermediate Representasion es la representación intermedia de bajo nivel, donde se aplican optimizaciones específicas de hardware.
  • Generación de código máquina: se emite código para LLVM, NVCC, o cualquier backend propietario a través del marco de generación de TVM.

Capa de algoritmo

Optimización de operadores

Antes de compilar es necesario verificar que cada operador del modelo sea soportado por el destino. Si no lo es, se puede reemplazar por otro funcionalmente similar. Un ejemplo clásico es YOLOv5, que usa SiLU por defecto; esta activación incluye exponencial y división, lo cual es costoso en FPGA. Reemplazar SiLU por ReLU6 mantiene la precisión y simplifica el hardware.

Otra línea de trabajo es reducir la complegidad de las capas convolucionales:

  • Reducción de canales: SqueezeNet usa filtros 1×1 para comprimir la entrada antes de expandirla.
  • Convolución separable en profundidad: emplea Depthwise + Pointwise, como en MobileNet y Xception.
  • Cuello de botella lineal: MobileNetV2 reemplaza la última ReLU por una activación lineal para evitar pérdida de información.
  • Convolución por grupos: divide los canales de entrada en grupos independientes, reduciendo las operaciones MAC.

Cuantización en PyTorch

PyTorch ofrece dos APIs principales: cuantización en modo eager y cuantización basada en torch.fx.

Cuantización eager

En este modo el usuario debe insertar manualmente los nodos QuantStub y DeQuantStub, declarar todos los operadores en __init__ y especificar los patrones de fusión. El siguiente ejemplo muestra una versión simplificada para ResNet-18:

import torch
import torchvision
import torch.nn as nn
import torch.optim as optim

dispositivo = torch.device("cuda" if torch.cuda.is_available() else "cpu")
base = torchvision.models.resnet18(weights=None)
base.load_state_dict(torch.load("./weights/resnet18.pth", map_location=dispositivo), strict=False)

def patrones_fusion_resnet18():
    patrones = [["modelo_fp32.conv1", "modelo_fp32.bn1", "modelo_fp32.relu"]]
    for etapa in range(1, 5):
        for bloque in range(2):
            patrones.append([
                f"modelo_fp32.layer{etapa}.{bloque}.conv1",
                f"modelo_fp32.layer{etapa}.{bloque}.bn1",
                f"modelo_fp32.layer{etapa}.{bloque}.relu",
            ])
            patrones.append([
                f"modelo_fp32.layer{etapa}.{bloque}.conv2",
                f"modelo_fp32.layer{etapa}.{bloque}.bn2",
            ])
    return patrones

class EnvolturaCuantizada(nn.Module):
    def __init__(self, modelo_fp32):
        super().__init__()
        self.cuant = torch.ao.quantization.QuantStub()
        self.descuant = torch.ao.quantization.DeQuantStub()
        self.modelo_fp32 = modelo_fp32

    def forward(self, x):
        x = self.cuant(x)
        x = self.modelo_fp32(x)
        return self.descuant(x)

modelo_qat = EnvolturaCuantizada(base)
modelo_qat.eval()
modelo_qat = torch.ao.quantization.fuse_modules(modelo_qat, patrones_fusion_resnet18())
modelo_qat.train()

modelo_qat.qconfig = torch.ao.quantization.QConfig(
    activation=torch.ao.quantization.MinMaxObserver.with_args(dtype=torch.quint8),
    weight=torch.ao.quantization.MinMaxObserver.with_args(dtype=torch.qint8,
                                                          qscheme=torch.per_tensor_affine),
)
torch.ao.quantization.prepare_qat(modelo_qat, inplace=True)
modelo_qat.to(dispositivo)

# Entrenamiento ...
optimizador = optim.Adam(modelo_qat.parameters(), lr=1e-4)

# Tras entrenar, convertir a INT8
modelo_qat.eval()
modelo_int8 = torch.ao.quantization.convert(modelo_qat, inplace=True)
torch.save(modelo_int8.state_dict(), "./weights/resnet18_int8.pth")

Cuantización FX

El modo FX traza el grafo completo, por lo que fusiona operadores y gestiona los stubs automáticamente. Solo hay que declarar las configuraciones por tipo o por nombre de capa:

import copy
from torch.ao.quantization import get_default_qat_qconfig
from torch.ao.quantization.quantize_fx import prepare_qat_fx, convert_fx

def crear_modelo_qat_fx(modelo_fp32):
    qconfig_global = get_default_qat_qconfig("qnnpack")
    qconfig_dict = {
        "": qconfig_global,
        "object_type": [
            (torch.nn.Softmax, qconfig_global),
            (torch.nn.Embedding, None),      # capa sin cuantizar
        ],
    }
    preparado = prepare_qat_fx(copy.deepcopy(modelo_fp32), qconfig_dict)
    return preparado

modelo_qat = crear_modelo_qat_fx(base)
# entrenamiento ...
modelo_cuantizado = convert_fx(modelo_qat)
torch.onnx.export(
    modelo_cuantizado,
    (entrada1, entrada2, entrada3),
    "modelo_qat.onnx",
    input_names=["e1", "e2", "e3"],
    output_names=["salida"],
    opset_version=16,
)

Capa de compilación

Parseo de modelos cuantizados con QNN

QNN es el dialecto de TVM para importar modelos ya cuantizados. Cada operador del framwork origen se mapea a un operador Relay mediante un convert_map. Por ejemplo, quantized::conv2d se traduce a una función como la siguiente:

def _conv2d_cuantizado(con_relu=False):
    def _implementar(entradas, _):
        # entradas[0]: activación de entrada
        # entradas[1]: (peso, escala_peso, zp_peso, sesgo)
        # entradas[2-5]: stride, padding, dilation, groups
        # entradas[6-7]: escala_salida, zp_salida
        # entradas[8-9]: escala_entrada, zp_entrada (añadidas por el frontend)
        params = entradas[1]
        peso = params[0]
        escala_peso = params[1]
        zp_peso = params[2]
        sesgo = params[3]

        if len(params) > 4:
            stride = params[4]
            padding = params[5]
            dilation = params[6]
            grupos = params[7]
            escala_salida = _expr.const(entradas[2])
            zp_salida = _expr.const(entradas[3])
            assert len(entradas) == 6, "Faltan parámetros de cuantización de entrada"
            escala_entrada = _expr.const(entradas[4])
            zp_entrada = _expr.const(entradas[5])
        else:
            stride = entradas[2]
            padding = entradas[3]
            dilation = entradas[4]
            grupos = entradas[5]
            escala_salida = _expr.const(entradas[6])
            zp_salida = _expr.const(entradas[7])
            assert len(entradas) == 10
            escala_entrada = _expr.const(entradas[8])
            zp_entrada = _expr.const(entradas[9])

        forma_peso = infer_shape(peso)
        kernel = (forma_peso[2], forma_peso[3])
        canales_salida = forma_peso[0]

        if padding[0] != 0 or padding[1] != 0:
            valor_relleno = _get_scalar(zp_entrada)
            activ = _op.nn.pad(
                entradas[0],
                pad_width=((0, 0), (0, 0),
                           (padding[0], padding[0]), (padding[1], padding[1])),
                pad_value=float(valor_relleno),
            )
        else:
            activ = entradas[0]

        salida_conv = relay.qnn.op.conv2d(
            activ,
            peso,
            input_zero_point=zp_entrada,
            kernel_zero_point=zp_peso,
            input_scale=escala_entrada,
            kernel_scale=escala_peso,
            kernel_size=kernel,
            strides=stride,
            padding=(0, 0),
            dilation=dilation,
            groups=grupos,
            channels=canales_salida,
        )

        return _requantizar_y_sumar_bias(
            salida_conv, sesgo, escala_entrada, escala_peso,
            escala_salida, zp_salida, con_relu
        )

    return _implementar

QNN expande un operador cuantizado en una secuenncia de operaciones Relay ya existentes. Si A y W son tensores cuantizados, el producto escalar se descompone en:

Σ QA·QW
- Σ zp_W·QA
- Σ zp_A·QW
+ N·zp_A·zp_W

De este modo se reutilizan los kernels existentes del compilador.

Cálculo por operadores

Multiplicación: un producto de enteros se transforma en una multiplicación entera seguida de una re-cuantización:

output = (input0 - zp0)*(input1 - zp1)*(scale0*scale1/scale_out) + zp_out

Re-cuantización: el factor escalarse descompone en un multiplicador entero y un desplazamiento, evitando la aritmética en coma flotante en tiempo de ejecución:

#include <cmath>
#include <limits>
#include <utility>

std::pair<int32_t, int32_t> descomponer_multiplicador(double factor) {
    int32_t exponente;
    if (factor == 0.0) return {0, 0};

    double mantisa = std::frexp(factor, &exponente);
    mantisa = std::round(mantisa * (1LL << 31));
    int64_t mantisa_i = static_cast<int64_t>(mantisa);

    if (mantisa_i == (1LL << 31)) {
        mantisa_i >>= 1;
        ++exponente;
    }
    return {static_cast<int32_t>(mantisa_i), exponente};
}

Suma: los dos operandos deben re-cuantizarse a la misma escala antes de sumarse, a diferencia del producto.

Activaciones complejas: para tanh, sigmoid o exp se suelen emplear aproximaciones polinómicas o tablas de búsqueda (LUT). Con int8 solo existen 256 entradas posibles, por lo que un LUT es eficiente:

import torch
import numpy as np

def generar_lut_tanh(valores, escala):
    flotantes = valores * escala              # de-cuantizar
    activados = flotantes.tanh()
    cuantizados = (activados / escala).round().clamp(-128, 127).to(torch.int8)
    return cuantizados

rango_bajo, rango_alto = -128, 127
muestras = torch.tensor(
    np.linspace(rango_bajo, rango_alto, rango_alto - rango_bajo + 1)
)
tabla_tanh = generar_lut_tanh(muestras, escala=0.2)

En tiempo de ejecución:

int8_t tanh_int8(int8_t x) {
    return TANH_LUT[static_cast<uint8_t>(x + 128)];
}

Operadores pasivos: padding, ReLU, Clip y MaxPool no cambian la escala, por lo que no requieren re-cuantización. En padding se rellena con el zero point; en ReLU se compara con el zero point en lugar de con 0.

Optimización de grafo

TVM aplica una serie de passes sobre Relay para optimizar el grafo. A continuación se resumen los más relevantes:

Pass Función
DeadCodeElimination Elimina expresiones no utilizadas.
FoldConstant Evalúa en tiempo de compilación operadores con entradas constantes.
FuseOps Fusiona operadores en operadores compuestos.
SimplifyInference Simplifica batch_norm, dropout, layer_norm, etc.
FastMath Reemplaza exp, erf, tanh y softmax por versiones aproximadas.
DynamicToStatic Convierte operadores dinámicos a estáticos cuando es posible.
EliminateCommonSubexpr Reemplaza subexpresiones idénticas por una sola variable.
CombineParallelConv2D Fusiona convoluciones paralelas que comparten entrada.
CombineParallelDense Fusiona capas densas paralelas.
AlterOpLayout Cambia el layout de datos para optimizar el acceso a memoria.
Legalize Reemplaza operadores por formas válidas para un backend, muy usado en QNN.
PartitionGraph Divide el grafo según anotaciones de aceleradores externos.

Generación de backend con BYOC

BYOC (Bring Your Own Codegen) permite a los proveedores de hardware integrar sus compiladores en TVM. El flujo es:

  1. Importar el modelo y representarlo en Relay.
  2. Aplicar optimizaciones independientes de hardware.
  3. Particionar el grafo en regiones que se ejecutarán en el host y en el acelerador.

La partición se realiza en la IR de alto nivel, ya que muchos aceleradores mapean directamente operadores Relay a kernels propietarios.

Anotación y agrupación por patrones

Se definen patrones como Conv2D + Add + ReLU para formar operadores compuestos. Luego se anotan los operadores soportados con un atributo de destino:

# Ejemplo conceptual
@tvm.ir.transform
def anotar_conv2d_para_accel(expr):
    if es_conv2d_soportada(expr):
        return expr.with_attr("target", "MiAcelerador")
    return expr

Partición basada en costo

No siempre es conveniente fusionar todo lo soportado; el tamaño del buffer local o el número de unidades de cómputo pueden limitar el tamaño de una región. BYOC permite establecer umbrales de costo y, si el traslado a un acelerador no compensa el overhead de transferencia, mantiene la región en el host.

Formatos de código generado

  • Representación JSON: fácil de interpretar por un motor de grafo; usado por TensorRT y Arm Compute Library.
  • C estándar: emite llamadas a kernels de una biblioteca propietaria y se enlaza con el módulo host.
  • Representación propietaria: para aceleradores con formatos propios, como Ethos-N o Vitis AI, se serializa un flujo de bits.

Runtime

El runtime ejecuta el grafo y distribuye subgrafos a cada backend:

  1. Módulo de metadatos: mantiene los pesos constantes y los buffers de entrada, salida e intermedios.
  2. Host: el motor de ejecución visita los nodos del grafo y usa los kernels generados para CPU/GPU.
  3. Acelerador: cuando se encuentra una función externa, se inicializa el motor del proveedor y se lanza el kernel sobre los buffers correspondientes.

Capa de simulación

Simulación de operadores

Antes de tocar hardware, se valida el compilador en varios niveles:

  1. Correctitud funcional: se escribe el operador en el DSL de TVM con un schedule por defecto y se compara contra PyTorch.
  2. Simulador del compilador: se genera el flujo de instrucciones del acelerador y se implementa un simulador en C++ que replica la lógica de cómputo, sincronización y memoria.
  3. Simulación de hardware: se alimentan DDR virtuales e instrucciones al simulador de C/RTL para verificar la correctitud del acelerador.
  4. Prueba en placa: una vez superadas las simulaciones, se ejecuta el modelo en el hardware real.

Simulación de redes

Tras validar operadores individuales, se construyen benchmarks de subredes que representan los patrones recurrentes del modelo completo. Cuando estas subredes pasan correctamente, se evalúa la red entera. Las métricas más útiles son:

  • Error relativo: comparación contra la inferencia en CPU.
  • Métrica de tarea: precisión top-1/top-5, accuracy, etc., sobre el conjunto de test con el pre y postprocesamiento reales.

Etiquetas: TVM Relay TIR AutoTVM QNN

Publicado el 9-16 10:52