Quantisierungsschicht laden: AWQ / GPTQ / Marlin / fp8
Verantwortung
vLLM behandelt Quantisierung nicht als nachträglichen "Trainieren-dann-Komprimieren"-Schritt, sondern implementiert jedes Quantisierungsschema als Paar aus QuantizationConfig + LinearMethodBase (sowie FusedMoEMethodBase für MoE). Das Config liest das Feld quantization_config aus dem Hugging-Face-Checkpoint und entscheidet, welches LinearMethod verwendet wird. Das LinearMethod übernimmt drei Aufgaben: create_weights legt zur Schichtkonstruktion die Parameterobjekte qweight / scales / qzeros an, process_weights_after_loading führt nach dem vollständigen Befüllen eine einmalige Umsortierung durch, und apply ruft im Forward den eigentlichen Quantisierungs-GEMM-Kernel auf. Diese Abstraktion erlaubt es, dasselbe ColumnParallelLinear sowohl unquantisiert als auch mit AWQ / GPTQ / fp8 / Marlin zu betreiben – nur das quant_method unterscheidet sich.
AWQ und GPTQ laufen in vLLM beide über die Marlin-Kernel-Familie: AutoAWQMarlinLinearMethod(auto_awq.py:414-426) und AutoGPTQLinearMethod(auto_gptq.py:306-324) wählen über choose_mp_linear_kernel den optimalen Kernel aus Conch / Exllama / Marlin / Machete aus. fp8 geht einen anderen Weg: Fp8Config(fp8.py:95-101) unterstützt Per-Tensor- / Per-Block- / statische / dynamische Aktivierungsquantisierung und wählt über init_fp8_linear_kernel das cutlass / Marlin / triton-Backend.
Entwurfsmotivation
Warum wird Quantisierung in Config + Method + Kernel aufgeteilt?
- Einheitliche Schnittstelle: Die Signatur von
LinearMethodBase.create_weights(layer, input_size_per_partition, output_partition_sizes, input_size, output_size, ...)(linear.py:144-168) ist für alle Quantisierungen gleich.ColumnParallelLinearmuss nicht wissen, ob es AWQ oder fp8 ist – es genügt, wennquant_config.get_quant_method(layer)eineLinearMethodBase-Instanz zurückgibt. - Kernel-Auswahl wird bis zur Laufzeit verzögert:
create_weightsvon AWQ / GPTQ konstruiert zuerst einMPLinearLayerConfigmit weight_type / group_size / zero_points / has_g_idx und wählt dann perchoose_mp_linear_kernel(...)einen konkreten Kernel(auto_gptq.py:341-354). Der Kernel selbst entscheidet, obprocess_weights_after_loadingeine Umsortierung vornehmen soll. - AWQ → GPTQ-Formatkonvertierung: AWQ-Checkpoints verwenden eine nicht standardgemäße Packing-Reihenfolge und packen entlang der Ausgabedimension; der Marlin-Kernel akzeptiert nur GPTQ-Stil (Packing entlang der Eingabedimension, standardmäßige Bit-Reihenfolge). Daher ruft
process_weights_after_loadingzuerst_convert_awq_to_standard_formatauf und übergibt dann an den Kernel(auto_awq.py:528-536). - GPTQs desc_act / g_idx:
AutoGPTQLinearMethod.create_weightslegt anhand vondesc_actfest, ob eing_idx-Parameter angelegt wird, und anhand vonmarlin_repeat_scales_on_all_ranks, ob die Scales unter TP repliziert oder geschhardet werden(auto_gptq.py:366-440). - fp8-Checkpoint nicht fp8 → Online-Quantisierung: Wenn der Checkpoint bf16 ist, ruft
process_weights_after_loadingprocess_fp8_weight_tensor_strategyauf, um ihn online zu fp8 zu quantisieren, und erledigt dabei auch das Per-Tensor-Reshape(fp8.py:416-437). - Marlin-Umsortierung: Wenn fp8
MarlinFP8ScaledMMLinearKernelwählt, transponiertprocess_weights_after_loadingdie Weight, schreibt sie nach der von Marlin erwarteten Anordnung neu und setztmarlin_input_dtype(fp8.py:398-408). - Parameter über v2-Loader: AWQ / GPTQ verwenden v2-Parameterklassen wie
PackedvLLMParameter/GroupQuantScaleParameter/PackedColumnParameter, diepacked_dim/packed_factor/input_dim/output_dimals Metadaten mitbringen, anhand dererweight_loader_v2korrekt schneidet(auto_awq.py:478-519).
Schlüsseldateien
QuantizationConfig:87-101— abstrakte Basisklasse, definiert Schnittstellen wieget_quant_method,get_cache_method.LinearMethodBase:141-180— Basisklasse der Quantisierungsmethode,create_weights/apply/process_weights_after_loading.AutoAWQConfig:171-180— parst diequantization_configdes AWQ-Checkpoints.AutoAWQMarlinLinearMethod.create_weights:439-526— legtqweight/qzeros/scalesan, wählt den Marlin-Kernel.AutoAWQMarlinLinearMethod.process_weights_after_loading:528-536—_convert_awq_to_standard_format+ Umsortierung des Kernels._convert_awq_to_standard_format:93-100— Konvertierungsfunktion AWQ-Packing → GPTQ-Packing.AutoGPTQConfig:97-103— parst den GPTQ-Checkpoint, behandeltbits/group_size/desc_act.AutoGPTQLinearMethod.create_weights:326-453— legtqweight/g_idx/scales/qzerosan, wählt anhand vondesc_actzwischenChannelQuantScaleParameterundGroupQuantScaleParameter.AutoGPTQLinearMethod.process+apply:455-464—kernel.process_weights_after_loading(layer)+kernel.apply_weights(layer, x, bias).Fp8Config:95-101— fp8-Konfiguration, unterscheidet Per-Tensor / Block / statische-dynamische Aktivierung.Fp8LinearMethod.create_weights:322-394— legtweight/weight_scale/ optionalinput_scalean, wählt überinit_fp8_linear_kernelden Kernel.Fp8LinearMethod.process_weights_after_loading:398-441— im Marlin-Fall Transposition + Kernel-Umsortierung, sonst Online-Quantisierung nicht-fp8-Checkpoints.choose_mp_linear_kernel:685-687— wählt nach Plattform / Form Conch / Exllama / Marlin / Machete usw.MarlinLinearKernel.process_weights_after_loading:88-137—ops.gptq_marlin_repacksortiert qweight in das von Marlin erwartete Layout um.
Datenfluss
Am Beispiel AWQ Marlin: ColumnParallelLinear.__init__ ruft self.quant_method.create_weights(...) auf. AutoAWQMarlinLinearMethod.create_weights konstruiert zuerst ein MPLinearLayerConfig, wählt per choose_mp_linear_kernel eine MarlinLinearKernel-Instanz und legt dann drei v2-Parameterobjekte an:
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) Danach schreibt der DefaultModelLoader über model.load_weights die qweight / scales / qzeros vom Datenträger über den v2-Loader in die Parameter. Nach dem vollständigen Befüllen ruft ColumnParallelLinear quant_method.process_weights_after_loading auf. AutoAWQMarlinLinearMethod wandelt zuerst das AWQ-Packing in den GPTQ-Stil um und übergibt an MarlinLinearKernel.process_weights_after_loading, das seinerseits mit ops.gptq_marlin_repack die Gewichte in das vom Marlin-Kernel erwartete Layout umsortiert(marlin.py:124-137). Im Forward ruft apply self.kernel.apply_weights(layer, x, bias) auf(auto_awq.py:538-545), was schließlich im Marlin CUDA- / Triton-Kernel landet.
Grenzen und Fehler
- Plattform unterstützt Marlin nicht:
AutoAWQMarlinLinearMethod.__init__ruft auf Nicht-CPU-Plattformenverify_marlin_supported(quant_type, group_size, has_zp=...)auf und wirft bei Nichterfüllung einen Fehler(auto_awq.py:431-437). - group_size=-1: In
create_weightswird group_size aufinput_sizezurückgebildet, also Per-Channel-Quantisierung(auto_awq.py:452-455). - fp8 block_quant + act_q_static: Block-Quantisierung erzwingt dynamische Aktivierung (
assert not self.act_q_static), sonst raise(fp8.py:367-368). - fp8 nicht-fp8-Checkpoints:
process_weights_after_loadingläuft überprocess_fp8_weight_tensor_tensor_strategy, um bf16 online zu quantisieren; bei Per-Tensor muss zusätzlich für fused Modules umgeformt werden(fp8.py:416-429). - GPTQ-Scales-Replikation bei Row-Parallel:
marlin_repeat_scales_on_all_ranksentscheidet, ob die Scales alsChannelQuantScaleParameter(scale_dim=None, auf allen Ranks repliziert) oder alsGroupQuantScaleParameter(scale_dim=0, geschshardet) angelegt werden(auto_gptq.py:367-440). - Phi-3 verschmolzenes QKV: Wenn QKV auf der Festplatte bereits verschmolzen ist, gibt es keine Shard-ID;
QKVParallelLinear._load_fused_module_from_checkpointsplittet selbst in drei Segmente(linear.py:1049-1097) – bei der Quantisierung ist dabei der Offset despacked_dimanzupassen. - Marlin act_type ist fp8:
MarlinLinearKernel.process_weights_after_loadingruft beic.act_type == torch.float8_e4m3fnops.marlin_int4_fp8_preprocessauf und multipliziert die Scales mit 512(marlin.py:98-103).
Zusammenfassung
Die Quantisierungs-Abstraktion trennt Config (Checkpoint-Metadaten lesen) → Method (Parameter anlegen + Nachverarbeitung + apply-Delegation) → Kernel (tatsächliches Rechnen) in drei Schichten. AWQ / GPTQ / Marlin teilen sich MPLinearLayerConfig + choose_mp_linear_kernel, fp8 geht über sein eigenes init_fp8_linear_kernel. Die Abstraktion verlässt sich auf tensorparallele lineare Schichten, um beim Anlegen der Parameter den v2-Loader einzuhängen; der Ladeablauf wird vom DefaultModelLoader angetrieben. Die nachverarbeiteten Gewichte werden im Forward schließlich von der Attention-Schicht und dem durch Blocktabelle beschriebenen KV-Cache konsumiert.