Skip to content

Carga de capas cuantizadas: AWQ / GPTQ / Marlin / fp8

源码版本v0.25.1

Responsabilidades

vLLM no trata la cuantización como un paso posterior de "comprimir tras el entrenamiento", sino que implementa cada esquema como un par QuantizationConfig + LinearMethodBase (y FusedMoEMethodBase para MoE). El Config lee el campo quantization_config del checkpoint de Hugging Face y decide qué LinearMethod usar; el LinearMethod se encarga de tres cosas: create_weights construye los objetos de parámetros qweight / scales / qzeros durante la creación de la capa, process_weights_after_loading hace un reordenamiento único tras volcar todos los pesos, y apply invoca al kernel GEMM cuantizado real en el forward. Esta abstracción permite que el mismo ColumnParallelLinear corra sin cuantizar o con AWQ / GPTQ / fp8 / Marlin, cambiando solo el quant_method.

AWQ y GPTQ en vLLM van por la familia de kernels Marlin: AutoAWQMarlinLinearMethod (auto_awq.py:414-426) y AutoGPTQLinearMethod (auto_gptq.py:306-324) eligen el óptimo entre Conch / Exllama / Marlin / Machete vía choose_mp_linear_kernel. fp8 va por otra ruta: Fp8Config (fp8.py:95-101) soporta cuantización per-tensor / per-block / con activaciones estáticas / dinámicas y selecciona entre backends cutlass / Marlin / triton mediante init_fp8_linear_kernel.

Motivación de diseño

¿Por qué partir la cuantización en tres capas Config + Method + Kernel?

  • Interfaz unificada: la firma LinearMethodBase.create_weights(layer, input_size_per_partition, output_partition_sizes, input_size, output_size, ...) (linear.py:144-168) es la misma para todos los esquemas; ColumnParallelLinear no necesita saber si es AWQ o fp8, solo que quant_config.get_quant_method(layer) devuelva una instancia de LinearMethodBase.
  • Selección del kernel se pospone a runtime: el create_weights de AWQ / GPTQ primero construye un MPLinearLayerConfig que describe weight_type / group_size / zero_points / has_g_idx y luego choose_mp_linear_kernel(...) elige un kernel concreto (auto_gptq.py:341-354); el propio kernel decide si necesita process_weights_after_loading para reordenar.
  • Conversión de formato AWQ → GPTQ: el checkpoint AWQ usa un orden de packing no estándar a lo largo de la dimensión de salida, mientras que el kernel Marlin solo acepta el estilo GPTQ (packing a lo largo de la dimensión de entrada, orden de bits estándar), por eso process_weights_after_loading primero llama a _convert_awq_to_standard_format y luego entrega el resultado al kernel (auto_awq.py:528-536).
  • desc_act / g_idx en GPTQ: AutoGPTQLinearMethod.create_weights decide según desc_act si crear el parámetro g_idx, y según marlin_repeat_scales_on_all_ranks si las scales se replican o se parten bajo TP (auto_gptq.py:366-440).
  • Cuantización online si el checkpoint no es fp8: si el checkpoint es bf16, process_weights_after_loading llama a process_fp8_weight_tensor_strategy para cuantizarlo a fp8 al vuelo y de paso normaliza per-tensor (fp8.py:416-437).
  • Reordenamiento Marlin: cuando fp8 elige MarlinFP8ScaledMMLinearKernel, process_weights_after_loading transpone el weight, lo reescribe en el layout que espera Marlin y fija marlin_input_dtype (fp8.py:398-408).
  • Los parámetros van por el cargador v2: AWQ / GPTQ usan clases v2 como PackedvLLMParameter / GroupQuantScaleParameter / PackedColumnParameter, que llevan metadatos packed_dim / packed_factor / input_dim / output_dim para que weight_loader_v2 los parta correctamente (auto_awq.py:478-519).

Archivos clave

Flujo de datos

Tomando AWQ Marlin como ejemplo: ColumnParallelLinear.__init__ llama a self.quant_method.create_weights(...), AutoAWQMarlinLinearMethod.create_weights primero construye un MPLinearLayerConfig, luego choose_mp_linear_kernel elige una instancia de MarlinLinearKernel y por último construye tres objetos de parámetros v2:

python
qweight = PackedvLLMParameter(
    data=torch.empty(
        input_size_per_partition,
        output_size_per_partition // self.quant_config.pack_factor,
        dtype=torch.int32,
    ),
    input_dim=0,
    output_dim=1,
    packed_dim=1,
    packed_factor=self.quant_config.pack_factor,
    weight_loader=weight_loader,
)

num_groups = input_size_per_partition // group_size

qzeros = PackedvLLMParameter(
    data=torch.empty(
        num_groups,
        output_size_per_partition // self.quant_config.pack_factor,
        dtype=torch.int32,
    ),
    input_dim=0,
    output_dim=1,
    packed_dim=1,
    packed_factor=self.quant_config.pack_factor,
    weight_loader=weight_loader,
)

scales = GroupQuantScaleParameter(
    data=torch.empty(
        num_groups,
        output_size_per_partition,
        dtype=params_dtype,
    ),
    input_dim=0,
    output_dim=1,
    weight_loader=weight_loader,
)

(auto_awq.py:478-515) Después DefaultModelLoader, a través de model.load_weights, escribe en disco qweight / scales / qzeros siguiendo el particionado del cargador v2. Una vez volcados todos, ColumnParallelLinear llama a quant_method.process_weights_after_loading; AutoAWQMarlinLinearMethod primero convierte el packing AWQ al estilo GPTQ y se lo pasa a MarlinLinearKernel.process_weights_after_loading, que con ops.gptq_marlin_repack reordena los pesos al layout que espera el kernel Marlin (marlin.py:124-137). En el forward, apply invoca self.kernel.apply_weights(layer, x, bias) (auto_awq.py:538-545) y termina cayendo en los kernels CUDA / Triton de Marlin.

Límites y fallos

  • Plataforma sin soporte Marlin: AutoAWQMarlinLinearMethod.__init__ en plataformas no CPU llama a verify_marlin_supported(quant_type, group_size, has_zp=...) y lanza un error si no se cumple (auto_awq.py:431-437).
  • group_size=-1: create_weights degenera group_size a input_size, es decir, cuantización per-channel (auto_awq.py:452-455).
  • fp8 block_quant + act_q_static: la cuantización por bloques exige activaciones dinámicas (assert not self.act_q_static); en caso contrario hace raise (fp8.py:367-368).
  • fp8 con checkpoint no fp8: process_weights_after_loading pasa por process_fp8_weight_tensor_tensor_strategy para cuantizar online bf16; además, en per-tensor hay que renormalizar para módulos fused (fp8.py:416-429).
  • Replicación de scales en GPTQ row-parallel: marlin_repeat_scales_on_all_ranks decide si las scales son ChannelQuantScaleParameter (scale_dim=None, replicadas en todos los ranks) o GroupQuantScaleParameter (scale_dim=0, partidas) (auto_gptq.py:367-440).
  • Phi-3 con QKV fusionado: cuando el QKV ya viene fusionado en disco no tiene shard id; QKVParallelLinear._load_fused_module_from_checkpoint lo parte en tres segmentos (linear.py:1049-1097) y, al cuantizar, hay que ajustar el offset del packed_dim.
  • act_type de Marlin es fp8: cuando c.act_type == torch.float8_e4m3fn, MarlinLinearKernel.process_weights_after_loading llama a ops.marlin_int4_fp8_preprocess y multiplica las scales por 512 (marlin.py:98-103).

Resumen

La abstracción de cuantización separa Config (lee metadatos del checkpoint) → Method (construye parámetros + post-procesado + delega en apply) → Kernel (el cálculo real) en tres capas; AWQ / GPTQ / Marlin comparten MPLinearLayerConfig + choose_mp_linear_kernel, mientras que fp8 va por su propio init_fp8_linear_kernel. Depende de capas lineales con tensor parallelism para que, al construir los pesos, el cargador v2 se inserte en los parámetros; el flujo de carga lo orquesta DefaultModelLoader, y los pesos ya post-procesados se consumen en el forward tanto en las capas de attention como en la KV cache descrita por la tabla de bloques.

Véase la documentación oficial: vLLM 文档 · README