Carga de capas cuantizadas: AWQ / GPTQ / Marlin / fp8
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;ColumnParallelLinearno necesita saber si es AWQ o fp8, solo quequant_config.get_quant_method(layer)devuelva una instancia deLinearMethodBase. - Selección del kernel se pospone a runtime: el
create_weightsde AWQ / GPTQ primero construye unMPLinearLayerConfigque describe weight_type / group_size / zero_points / has_g_idx y luegochoose_mp_linear_kernel(...)elige un kernel concreto (auto_gptq.py:341-354); el propio kernel decide si necesitaprocess_weights_after_loadingpara 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_loadingprimero llama a_convert_awq_to_standard_formaty luego entrega el resultado al kernel (auto_awq.py:528-536). desc_act/g_idxen GPTQ:AutoGPTQLinearMethod.create_weightsdecide segúndesc_actsi crear el parámetrog_idx, y segúnmarlin_repeat_scales_on_all_rankssi 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_loadingllama aprocess_fp8_weight_tensor_strategypara cuantizarlo a fp8 al vuelo y de paso normaliza per-tensor (fp8.py:416-437). - Reordenamiento Marlin: cuando fp8 elige
MarlinFP8ScaledMMLinearKernel,process_weights_after_loadingtranspone el weight, lo reescribe en el layout que espera Marlin y fijamarlin_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 metadatospacked_dim/packed_factor/input_dim/output_dimpara queweight_loader_v2los parta correctamente (auto_awq.py:478-519).
Archivos clave
QuantizationConfig:87-101— clase abstracta base; defineget_quant_method,get_cache_method, etc.LinearMethodBase:141-180— clase base de métodos de cuantización;create_weights/apply/process_weights_after_loading.AutoAWQConfig:171-180— parsea elquantization_configdel checkpoint AWQ.AutoAWQMarlinLinearMethod.create_weights:439-526— construyeqweight/qzeros/scalesy elige el kernel Marlin.AutoAWQMarlinLinearMethod.process_weights_after_loading:528-536—_convert_awq_to_standard_format+ reordenamiento del propio kernel._convert_awq_to_standard_format:93-100— función de conversión de packing AWQ → packing GPTQ.AutoGPTQConfig:97-103— parsea el checkpoint GPTQ, gestionabits/group_size/desc_act.AutoGPTQLinearMethod.create_weights:326-453— construyeqweight/g_idx/scales/qzerosy segúndesc_acteligeChannelQuantScaleParameteroGroupQuantScaleParameter.AutoGPTQLinearMethod.process+apply:455-464—kernel.process_weights_after_loading(layer)+kernel.apply_weights(layer, x, bias).Fp8Config:95-101— configuración fp8; distingue per-tensor / block / activaciones estática-dinámica.Fp8LinearMethod.create_weights:322-394— construyeweight/weight_scaley opcionalinput_scale, y elige el kernel víainit_fp8_linear_kernel.Fp8LinearMethod.process_weights_after_loading:398-441— con Marlin transpone + reordena el kernel; en caso contrario cuantiza online los checkpoints no fp8.choose_mp_linear_kernel:685-687— según plataforma / forma elige Conch / Exllama / Marlin / Machete, etc.MarlinLinearKernel.process_weights_after_loading:88-137—ops.gptq_marlin_repackreordena qweight al layout que espera Marlin.
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:
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 averify_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_weightsdegeneragroup_sizeainput_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_loadingpasa porprocess_fp8_weight_tensor_tensor_strategypara 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_ranksdecide si las scales sonChannelQuantScaleParameter(scale_dim=None, replicadas en todos los ranks) oGroupQuantScaleParameter(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_checkpointlo parte en tres segmentos (linear.py:1049-1097) y, al cuantizar, hay que ajustar el offset delpacked_dim. act_typede Marlin es fp8: cuandoc.act_type == torch.float8_e4m3fn,MarlinLinearKernel.process_weights_after_loadingllama aops.marlin_int4_fp8_preprocessy 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.