Skip to content

Quantisierungsschicht laden: AWQ / GPTQ / Marlin / fp8

源码版本v0.25.1

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. ColumnParallelLinear muss nicht wissen, ob es AWQ oder fp8 ist – es genügt, wenn quant_config.get_quant_method(layer) eine LinearMethodBase-Instanz zurückgibt.
  • Kernel-Auswahl wird bis zur Laufzeit verzögert: create_weights von AWQ / GPTQ konstruiert zuerst ein MPLinearLayerConfig mit weight_type / group_size / zero_points / has_g_idx und wählt dann per choose_mp_linear_kernel(...) einen konkreten Kernel(auto_gptq.py:341-354). Der Kernel selbst entscheidet, ob process_weights_after_loading eine 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_loading zuerst _convert_awq_to_standard_format auf und übergibt dann an den Kernel(auto_awq.py:528-536).
  • GPTQs desc_act / g_idx: AutoGPTQLinearMethod.create_weights legt anhand von desc_act fest, ob ein g_idx-Parameter angelegt wird, und anhand von marlin_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_loading process_fp8_weight_tensor_strategy auf, um ihn online zu fp8 zu quantisieren, und erledigt dabei auch das Per-Tensor-Reshape(fp8.py:416-437).
  • Marlin-Umsortierung: Wenn fp8 MarlinFP8ScaledMMLinearKernel wählt, transponiert process_weights_after_loading die Weight, schreibt sie nach der von Marlin erwarteten Anordnung neu und setzt marlin_input_dtype(fp8.py:398-408).
  • Parameter über v2-Loader: AWQ / GPTQ verwenden v2-Parameterklassen wie PackedvLLMParameter / GroupQuantScaleParameter / PackedColumnParameter, die packed_dim / packed_factor / input_dim / output_dim als Metadaten mitbringen, anhand derer weight_loader_v2 korrekt schneidet(auto_awq.py:478-519).

Schlüsseldateien

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:

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) 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-Plattformen verify_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_weights wird group_size auf input_size zurü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_loading läuft über process_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_ranks entscheidet, ob die Scales als ChannelQuantScaleParameter (scale_dim=None, auf allen Ranks repliziert) oder als GroupQuantScaleParameter (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_checkpoint splittet selbst in drei Segmente(linear.py:1049-1097) – bei der Quantisierung ist dabei der Offset des packed_dim anzupassen.
  • Marlin act_type ist fp8: MarlinLinearKernel.process_weights_after_loading ruft bei c.act_type == torch.float8_e4m3fn ops.marlin_int4_fp8_preprocess auf 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.

Siehe offizielle Dokumentation: vLLM 文档 · README.