Chargement des couches quantifiées : AWQ / GPTQ / Marlin / fp8
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 ;ColumnParallelLinearn'a pas besoin de savoir s'il s'agit d'AWQ ou de fp8, il suffit quequant_config.get_quant_method(layer)retourne une instance deLinearMethodBase. - Sélection du kernel repoussée au runtime : le
create_weightsde AWQ / GPTQ construit d'abord unMPLinearLayerConfigdécrivant weight_type / group_size / zero_points / has_g_idx, puischoose_mp_linear_kernel(...)choisit un kernel concret(auto_gptq.py:341-354), qui décide lui-même s'il faut réarranger viaprocess_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_loadingappelle donc_convert_awq_to_standard_formatavant de handed au kernel(auto_awq.py:528-536). - desc_act / g_idx de GPTQ :
AutoGPTQLinearMethod.create_weightsdécide selondesc_actde créer ou non le paramètreg_idx, et selonmarlin_repeat_scales_on_all_rankssi 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_loadingappelleprocess_fp8_weight_tensor_strategypour 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_loadingtranspose le weight, le réécrit selon la disposition attendue par Marlin et fixemarlin_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éespacked_dim/packed_factor/input_dim/output_dim;weight_loader_v2s'appuie dessus pour splitter correctement(auto_awq.py:478-519).
Fichiers clés
QuantizationConfig:87-101— classe de base abstraite, définitget_quant_method,get_cache_method, etc.LinearMethodBase:141-180— classe de base des méthodes de quantification :create_weights/apply/process_weights_after_loading.AutoAWQConfig:171-180— parse lequantization_configd'un checkpoint AWQ.AutoAWQMarlinLinearMethod.create_weights:439-526— construitqweight/qzeros/scales, sélectionne le kernel Marlin.AutoAWQMarlinLinearMethod.process_weights_after_loading:528-536—_convert_awq_to_standard_format+ réarrangement propre au kernel._convert_awq_to_standard_format:93-100— fonction de conversion packing AWQ → packing GPTQ.AutoGPTQConfig:97-103— parse un checkpoint GPTQ, gèrebits/group_size/desc_act.AutoGPTQLinearMethod.create_weights:326-453— construitqweight/g_idx/scales/qzeros, choisitChannelQuantScaleParameterouGroupQuantScaleParameterselondesc_act.AutoGPTQLinearMethod.process+apply:455-464—kernel.process_weights_after_loading(layer)+kernel.apply_weights(layer, x, bias).Fp8Config:95-101— configuration fp8, distingue per-tensor / block / activations statique/dynamique.Fp8LinearMethod.create_weights:322-394— construitweight/weight_scale/input_scaleoptionnel, sélectionne le kernel viainit_fp8_linear_kernel.Fp8LinearMethod.process_weights_after_loading:398-441— transpose + réarrange pour Marlin, sinon quantifie en ligne un checkpoint non fp8.choose_mp_linear_kernel:685-687— sélectionne un kernel Conch / Exllama / Marlin / Machete selon la plateforme et la forme.MarlinLinearKernel.process_weights_after_loading:88-137—ops.gptq_marlin_repackréarrange qweight dans la disposition attendue par Marlin.
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 :
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__appelleverify_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_weightsdégrade group_size eninput_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_loadingpasse parprocess_fp8_weight_tensor_tensor_strategypour 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_ranksdétermine si les scales sont unChannelQuantScaleParameter(scale_dim=None, répliqué sur tous les ranks) ou unGroupQuantScaleParameter(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_checkpointdécoupe lui-même les trois segments(linear.py:1049-1097), et doit gérer l'ajustement d'offset depacked_dimcôté quantification. - Marlin act_type fp8 :
MarlinLinearKernel.process_weights_after_loadingappelleops.marlin_int4_fp8_preprocessquandc.act_type == torch.float8_e4m3fnet 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