-
-
Notifications
You must be signed in to change notification settings - Fork 20.4k
[FP8] Add opt-in ParallelLMHead dispatch to Fp8Config #41000
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from 1 commit
72238d4
3520a71
5ca163e
511412b
49a8a6b
de77662
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -38,6 +38,10 @@ | |
| LinearMethodBase, | ||
| UnquantizedLinearMethod, | ||
| ) | ||
| from vllm.model_executor.layers.vocab_parallel_embedding import ( | ||
| ParallelLMHead, | ||
| UnquantizedEmbeddingMethod, | ||
| ) | ||
| from vllm.model_executor.layers.quantization import QuantizationMethods | ||
| from vllm.model_executor.layers.quantization.base_config import ( | ||
| QuantizationConfig, | ||
|
|
@@ -102,10 +106,12 @@ def __init__( | |
| activation_scheme: str = "dynamic", | ||
| ignored_layers: list[str] | None = None, | ||
| weight_block_size: list[int] | None = None, | ||
| lm_head_quantized: bool = False, | ||
| ) -> None: | ||
| super().__init__() | ||
|
|
||
| self.is_checkpoint_fp8_serialized = is_checkpoint_fp8_serialized | ||
| self.lm_head_quantized = lm_head_quantized | ||
|
|
||
| if activation_scheme not in ACTIVATION_SCHEMES: | ||
| raise ValueError(f"Unsupported activation scheme {activation_scheme}") | ||
|
|
@@ -162,22 +168,29 @@ def from_config(cls, config: dict[str, Any]) -> "Fp8Config": | |
| ignored_layers = cls.get_from_keys_or( | ||
| config, ["modules_to_not_convert"], None | ||
| ) | ||
| lm_head_quantized = cls.get_from_keys_or(config, ["lm_head"], default=False) | ||
| return cls( | ||
| is_checkpoint_fp8_serialized=is_checkpoint_fp8_serialized, | ||
| activation_scheme=activation_scheme, | ||
| ignored_layers=ignored_layers, | ||
| weight_block_size=weight_block_size, | ||
| lm_head_quantized=lm_head_quantized, | ||
| ) | ||
|
|
||
| def get_quant_method( | ||
| self, layer: torch.nn.Module, prefix: str | ||
| ) -> "QuantizeMethodBase | None": | ||
| if isinstance(layer, LinearBase): | ||
| is_parallel_lm_head = isinstance(layer, ParallelLMHead) | ||
| if isinstance(layer, LinearBase) or ( | ||
| is_parallel_lm_head and self.lm_head_quantized | ||
| ): | ||
|
Comment on lines
+183
to
+186
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. The current implementation of Furthermore, as noted in the PR description, |
||
| if is_layer_skipped( | ||
| prefix=prefix, | ||
| ignored_layers=self.ignored_layers, | ||
| fused_mapping=self.packed_modules_mapping, | ||
| ): | ||
| if is_parallel_lm_head: | ||
| return UnquantizedEmbeddingMethod() | ||
| return UnquantizedLinearMethod() | ||
| if not self.is_checkpoint_fp8_serialized: | ||
| online_method = Fp8OnlineLinearMethod(self) | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Importing
ParallelLMHeadandUnquantizedEmbeddingMethodat the top level offp8.pyfromvllm.model_executor.layers.vocab_parallel_embeddingmay lead to circular import issues in the future, as quantization configs are often imported by the layers they configure. It is generally safer to perform these imports insideget_quant_methodor useTYPE_CHECKINGfor type hints andimportlibfor runtime checks if necessary.