Skip to content

Chargement des couches quantifiées : AWQ / GPTQ / Marlin / fp8

源码版本v0.25.1

Responsabilités

vLLM ne traite pas la quantification comme un post-traitement « on compresse après l'entraînement », mais implémente chaque schéma de quantification comme un couple QuantizationConfig + LinearMethodBase (et FusedMoEMethodBase pour les MoE). Le Config lit le champ quantization_config du checkpoint Hugging Face et décide quel LinearMethod utiliser ; le LinearMethod a trois responsabilités : create_weights construit les paramètres qweight / scales / qzeros à l'initialisation de la couche, process_weights_after_loading effectue un réarrangement one-shot une fois tous les poids chargés, et apply invoque le véritable kernel GEMM quantifié lors du forward. Cette abstraction permet au même ColumnParallelLinear de fonctionner en non quantifié ou en AWQ / GPTQ / fp8 / Marlin, seul le quant_method change.

AWQ et GPTQ utilisent dans vLLM la famille de kernels Marlin : AutoAWQMarlinLinearMethod(auto_awq.py:414-426) et AutoGPTQLinearMethod(auto_gptq.py:306-324), via choose_mp_linear_kernel qui choisit entre Conch / Exllama / Marlin / Machete. fp8 emprunte une autre voie : Fp8Config(fp8.py:95-101) supporte la quantification per-tensor / per-block / statique / dynamique des activations, et sélectionne un backend cutlass / Marlin / triton via init_fp8_linear_kernel.

Motivation de conception

Pourquoi séparer la quantification en trois couches Config + Method + Kernel ?

  • Interface unifiée : la signature LinearMethodBase.create_weights(layer, input_size_per_partition, output_partition_sizes, input_size, output_size, ...)(linear.py:144-168) est identique pour toutes les quantifications ; ColumnParallelLinear n'a pas besoin de savoir s'il s'agit d'AWQ ou de fp8, il suffit que quant_config.get_quant_method(layer) retourne une instance de LinearMethodBase.
  • Sélection du kernel repoussée au runtime : le create_weights de AWQ / GPTQ construit d'abord un MPLinearLayerConfig décrivant weight_type / group_size / zero_points / has_g_idx, puis choose_mp_linear_kernel(...) choisit un kernel concret(auto_gptq.py:341-354), qui décide lui-même s'il faut réarranger via process_weights_after_loading.
  • Conversion AWQ → format GPTQ : les checkpoints AWQ utilisent un ordre de packing non standard et empaquettent selon la dimension de sortie, alors que le kernel Marlin n'accepte que le style GPTQ (empaquetage selon la dimension d'entrée, ordre de bits standard). process_weights_after_loading appelle donc _convert_awq_to_standard_format avant de handed au kernel(auto_awq.py:528-536).
  • desc_act / g_idx de GPTQ : AutoGPTQLinearMethod.create_weights décide selon desc_act de créer ou non le paramètre g_idx, et selon marlin_repeat_scales_on_all_ranks si les scales sont répliquées ou shardées sous TP(auto_gptq.py:366-440).
  • Quantification fp8 en ligne quand le checkpoint n'est pas fp8 : si le checkpoint est en bf16, process_weights_after_loading appelle process_fp8_weight_tensor_strategy pour le quantifier en fp8 à la volée, en gérant au passage le remaniement per-tensor(fp8.py:416-437).
  • Réarrangement Marlin : quand fp8 sélectionne MarlinFP8ScaledMMLinearKernel, process_weights_after_loading transpose le weight, le réécrit selon la disposition attendue par Marlin et fixe marlin_input_dtype(fp8.py:398-408).
  • Paramètres via le loader v2 : AWQ / GPTQ utilisent les classes de paramètres v2 comme PackedvLLMParameter / GroupQuantScaleParameter / PackedColumnParameter, qui portent leurs métadonnées packed_dim / packed_factor / input_dim / output_dim ; weight_loader_v2 s'appuie dessus pour splitter correctement(auto_awq.py:478-519).

Fichiers clés

Flux de données

Avec AWQ Marlin pour exemple, ColumnParallelLinear.__init__ appelle self.quant_method.create_weights(...), AutoAWQMarlinLinearMethod.create_weights construit d'abord un MPLinearLayerConfig puis choose_mp_linear_kernel, qui retourne une instance de MarlinLinearKernel, avant de créer trois paramètres 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) Ensuite DefaultModelLoader via model.load_weights découpe et écrit qweight / scales / qzeros du disque selon le loader v2. Une fois tous les poids chargés, ColumnParallelLinear appelle quant_method.process_weights_after_loading, et AutoAWQMarlinLinearMethod convertit d'abord le packing AWQ en style GPTQ puis délègue à MarlinLinearKernel.process_weights_after_loading, qui utilise ops.gptq_marlin_repack pour réarranger les poids dans la disposition attendue par le kernel Marlin(marlin.py:124-137). Au forward, apply appelle self.kernel.apply_weights(layer, x, bias)(auto_awq.py:538-545), qui finit dans les kernels CUDA / Triton de Marlin.

Limites et échecs

  • Marlin non supporté sur la plateforme : AutoAWQMarlinLinearMethod.__init__ appelle verify_marlan_supported(quant_type, group_size, has_zp=...) sur plateforme non CPU, et lève une erreur si les conditions ne sont pas remplies(auto_awq.py:431-437).
  • group_size=-1 : create_weights dégrade group_size en input_size, c'est-à-dire une quantification per-channel(auto_awq.py:452-455).
  • fp8 block_quant + act_q_static : la quantification par blocs impose des activations dynamiques (assert not self.act_q_static), sinon raise(fp8.py:367-368).
  • Checkpoint fp8 non fp8 : process_weights_after_loading passe par process_fp8_weight_tensor_tensor_strategy pour quantifier le bf16 en ligne ; en per-tensor, il faut en plus remanier les modules fused(fp8.py:416-429).
  • Réplique des scales GPTQ en row-parallel : marlin_repeat_scales_on_all_ranks détermine si les scales sont un ChannelQuantScaleParameter (scale_dim=None, répliqué sur tous les ranks) ou un GroupQuantScaleParameter (scale_dim=0, sharded)(auto_gptq.py:367-440).
  • Phi-3 QKV fused : sur disque, quand QKV est déjà fused il n'y a pas de shard id ; QKVParallelLinear._load_fused_module_from_checkpoint découpe lui-même les trois segments(linear.py:1049-1097), et doit gérer l'ajustement d'offset de packed_dim côté quantification.
  • Marlin act_type fp8 : MarlinLinearKernel.process_weights_after_loading appelle ops.marlin_int4_fp8_preprocess quand c.act_type == torch.float8_e4m3fn et multiplie les scales par 512(marlin.py:98-103).

Résumé

L'abstraction de quantification sépare trois couches : Config (lit les métadonnées du checkpoint) → Method (crée les paramètres + post-traitement + délègue à apply) → Kernel (le calcul effectif). AWQ / GPTQ / Marlin partagent MPLinearLayerConfig + choose_mp_linear_kernel, fp8 passe par son propre init_fp8_linear_kernel. Elle s'appuie sur couches linéaires en parallélisme de tenseurs pour injecter le loader v2 dans les paramètres lors de create_weights ; le chargement est piloté par DefaultModelLoader. Les poids post-traités sont finalement consommés au forward par les couches d'attention et le cache KV décrit par block-table.

Voir la documentation officielle : Documentation vLLM · README