forked from Karylab-cklius/vllm
[Quantization] add online fp8 ptpc (#44132)
Signed-off-by: walterbm <walter.beller.morales@gmail.com> Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
parent
ac3409d162
commit
753e9d55e6
@@ -0,0 +1,79 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
"""E2E tests for online FP8 per-channel quantization.
|
||||
|
||||
Loads a BF16 model with ``--quantization fp8_per_channel`` (online
|
||||
quantization) and compares log-probabilities against the same model served in
|
||||
BF16 without quantization. This exercises the full pipeline: config parsing,
|
||||
``Fp8PtpcOnlineLinearMethod``, ``Fp8PtpcOnlineMoEMethod``, weight
|
||||
loading, online quantization / shuffling, and inference.
|
||||
|
||||
``example_prompts`` is a pytest fixture (from conftest.py) that loads 8
|
||||
diverse prompts from ``tests/prompts/example.txt``.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from tests.quantization.utils import is_quant_method_supported
|
||||
|
||||
from ..utils import check_logprobs_close
|
||||
|
||||
# Small MoE model that fits on a single GPU and exercises both linear + MoE.
|
||||
MOE_MODEL = "allenai/OLMoE-1B-7B-0125-Instruct"
|
||||
# Small dense model (no MoE) to validate the linear-only path.
|
||||
DENSE_MODEL = "Qwen/Qwen3-0.6B"
|
||||
|
||||
MAX_MODEL_LEN = 1024
|
||||
MAX_TOKENS = 4
|
||||
NUM_LOG_PROBS = 8
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not is_quant_method_supported("fp8"),
|
||||
reason="fp8 is not supported on this GPU type.",
|
||||
)
|
||||
@pytest.mark.quant_model
|
||||
@pytest.mark.parametrize("model", [DENSE_MODEL, MOE_MODEL], ids=["dense", "moe"])
|
||||
def test_fp8_per_channel_logprobs(
|
||||
vllm_runner,
|
||||
example_prompts,
|
||||
model: str,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Compare BF16 baseline logprobs against online per-channel-quantized
|
||||
model.
|
||||
|
||||
Runs the same model twice -- once in BF16 (baseline) and once with online
|
||||
FP8 per-channel quantization -- then checks that the top log-probabilities
|
||||
are close. Only 4 tokens are generated to keep the test fast while still
|
||||
catching numerical divergence beyond expected per-channel error.
|
||||
"""
|
||||
with monkeypatch.context() as m:
|
||||
m.setenv("TOKENIZERS_PARALLELISM", "true")
|
||||
|
||||
with vllm_runner(
|
||||
model,
|
||||
max_model_len=MAX_MODEL_LEN,
|
||||
enforce_eager=True,
|
||||
) as vllm_model:
|
||||
baseline_outputs = vllm_model.generate_greedy_logprobs(
|
||||
example_prompts, MAX_TOKENS, NUM_LOG_PROBS
|
||||
)
|
||||
|
||||
with vllm_runner(
|
||||
model,
|
||||
max_model_len=MAX_MODEL_LEN,
|
||||
enforce_eager=True,
|
||||
quantization="fp8_per_channel",
|
||||
) as vllm_model:
|
||||
test_outputs = vllm_model.generate_greedy_logprobs(
|
||||
example_prompts, MAX_TOKENS, NUM_LOG_PROBS
|
||||
)
|
||||
|
||||
check_logprobs_close(
|
||||
outputs_0_lst=baseline_outputs,
|
||||
outputs_1_lst=test_outputs,
|
||||
name_0="bf16",
|
||||
name_1="fp8_per_channel",
|
||||
)
|
||||
@@ -0,0 +1,103 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Tests for FP8 per-channel online quantization.
|
||||
|
||||
Per-output-channel weight scale + dynamic per-token activation scale.
|
||||
bf16/fp16 checkpoints are quantized at load time with one fp32 scale per
|
||||
output channel for weights and one fp32 scale per token for activations
|
||||
(computed dynamically inside the kernel). Run via
|
||||
`pytest tests/quantization/test_fp8_per_channel.py --forked`.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from tests.quantization.utils import is_quant_method_supported
|
||||
from vllm import _custom_ops as ops
|
||||
from vllm.config.quantization import (
|
||||
_ONLINE_SHORTHANDS,
|
||||
QUANT_KEY_NAMES,
|
||||
QuantizationConfigArgs,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.online.base import (
|
||||
_ONLINE_LINEAR_METHODS,
|
||||
_ONLINE_MOE_METHODS,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.online.fp8 import (
|
||||
Fp8PtpcOnlineLinearMethod,
|
||||
Fp8PtpcOnlineMoEMethod,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
kFp8StaticChannelSym,
|
||||
)
|
||||
from vllm.platforms import current_platform
|
||||
|
||||
|
||||
def test_fp8_per_channel_shorthand_registered() -> None:
|
||||
"""The `fp8_per_channel` CLI shorthand must resolve to a config that
|
||||
dispatches the per-channel methods. Guards against regressions in
|
||||
`_ONLINE_SHORTHANDS` / `_ONLINE_LINEAR_METHODS` / `_ONLINE_MOE_METHODS`
|
||||
drifting out of sync.
|
||||
"""
|
||||
args = _ONLINE_SHORTHANDS["fp8_per_channel"]
|
||||
assert isinstance(args, QuantizationConfigArgs)
|
||||
assert args.linear is not None
|
||||
assert args.moe is not None
|
||||
assert args.linear.weight is kFp8StaticChannelSym
|
||||
assert args.moe.weight is kFp8StaticChannelSym
|
||||
|
||||
assert _ONLINE_LINEAR_METHODS[kFp8StaticChannelSym] is Fp8PtpcOnlineLinearMethod
|
||||
assert _ONLINE_MOE_METHODS[kFp8StaticChannelSym] is Fp8PtpcOnlineMoEMethod
|
||||
|
||||
assert QUANT_KEY_NAMES["fp8_per_channel_static"] is kFp8StaticChannelSym
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not is_quant_method_supported("fp8"),
|
||||
reason="FP8 is not supported on this GPU type.",
|
||||
)
|
||||
def test_scaled_fp8_quant_per_channel_shape() -> None:
|
||||
"""Verify the kernel call per-channel quant depends on: passing a 2D
|
||||
weight to `ops.scaled_fp8_quant` with `use_per_token_if_dynamic=True`
|
||||
yields one scale per output row -- a [N, 1] fp32 tensor.
|
||||
"""
|
||||
x = (torch.randn(size=(96, 256), device="cuda") * 13).to(torch.bfloat16)
|
||||
y, s = ops.scaled_fp8_quant(x, scale=None, use_per_token_if_dynamic=True)
|
||||
assert y.shape == (96, 256)
|
||||
assert y.dtype == current_platform.fp8_dtype()
|
||||
assert s.shape == (96, 1)
|
||||
assert s.dtype == torch.float32
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not is_quant_method_supported("fp8"),
|
||||
reason="FP8 is not supported on this GPU type.",
|
||||
)
|
||||
def test_fp8_per_channel_online_quantization(
|
||||
vllm_runner,
|
||||
monkeypatch,
|
||||
) -> None:
|
||||
"""End-to-end smoke: load `facebook/opt-125m` bf16 with
|
||||
`quantization='fp8_per_channel'`, check a dense Linear is wrapped by
|
||||
`Fp8PtpcOnlineLinearMethod`, its weights are fp8 with per-channel
|
||||
scales (shape `[N, 1]`), and a short greedy generation works.
|
||||
"""
|
||||
monkeypatch.setenv("VLLM_ALLOW_INSECURE_SERIALIZATION", "1")
|
||||
|
||||
with vllm_runner(
|
||||
"facebook/opt-125m",
|
||||
quantization="fp8_per_channel",
|
||||
enforce_eager=True,
|
||||
) as llm:
|
||||
|
||||
def check_model(model):
|
||||
fc1 = model.model.decoder.layers[0].fc1
|
||||
assert isinstance(fc1.quant_method, Fp8PtpcOnlineLinearMethod)
|
||||
assert fc1.weight.dtype == current_platform.fp8_dtype()
|
||||
assert fc1.weight_scale.ndim == 2
|
||||
assert fc1.weight_scale.shape[-1] == 1
|
||||
assert fc1.input_scale is None
|
||||
|
||||
llm.apply_model(check_model)
|
||||
outputs = llm.generate_greedy(["Hello my name is"], max_tokens=4)
|
||||
print(outputs[0][1])
|
||||
@@ -13,6 +13,7 @@ from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
kFp8DynamicTensorSym,
|
||||
kFp8DynamicTokenSym,
|
||||
kFp8Static128BlockSym,
|
||||
kFp8StaticChannelSym,
|
||||
kFp8StaticTensorSym,
|
||||
kInt8StaticChannelSym,
|
||||
kMxfp4Dynamic,
|
||||
@@ -24,6 +25,7 @@ QUANT_KEY_NAMES: dict[str, QuantKey] = {
|
||||
"fp8_per_tensor_static": kFp8StaticTensorSym,
|
||||
"fp8_per_tensor_dynamic": kFp8DynamicTensorSym,
|
||||
"fp8_per_token": kFp8DynamicTokenSym,
|
||||
"fp8_per_channel_static": kFp8StaticChannelSym,
|
||||
"fp8_per_block_static": kFp8Static128BlockSym,
|
||||
"fp8_per_block_dynamic": kFp8Dynamic128Sym,
|
||||
"mxfp8": kMxfp8Dynamic,
|
||||
@@ -118,6 +120,12 @@ _ONLINE_SHORTHANDS: dict[str, QuantizationConfigArgs] = {
|
||||
linear=QuantSpec(weight=kFp8Static128BlockSym),
|
||||
moe=QuantSpec(weight=kFp8Static128BlockSym),
|
||||
),
|
||||
# Per-output-channel weight scale + dynamic per-token activation.
|
||||
# Same shape as llmcompressor's FP8_DYNAMIC recipe.
|
||||
"fp8_per_channel": QuantizationConfigArgs(
|
||||
linear=QuantSpec(weight=kFp8StaticChannelSym),
|
||||
moe=QuantSpec(weight=kFp8StaticChannelSym),
|
||||
),
|
||||
"mxfp8": QuantizationConfigArgs(
|
||||
linear=QuantSpec(weight=kMxfp8Dynamic),
|
||||
moe=QuantSpec(weight=kMxfp8Dynamic),
|
||||
|
||||
@@ -40,6 +40,7 @@ QuantizationMethods = Literal[
|
||||
# _ONLINE_SHORTHANDS by the assertion in get_quantization_config().
|
||||
"fp8_per_tensor",
|
||||
"fp8_per_block",
|
||||
"fp8_per_channel",
|
||||
"int8_per_channel_weight_only",
|
||||
"mxfp8",
|
||||
]
|
||||
|
||||
@@ -30,6 +30,8 @@ from vllm.model_executor.layers.quantization.online.fp8 import (
|
||||
Fp8PerBlockOnlineMoEMethod,
|
||||
Fp8PerTensorOnlineLinearMethod,
|
||||
Fp8PerTensorOnlineMoEMethod,
|
||||
Fp8PtpcOnlineLinearMethod,
|
||||
Fp8PtpcOnlineMoEMethod,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.online.int8 import (
|
||||
Int8OnlineMoEMethod,
|
||||
@@ -41,6 +43,7 @@ from vllm.model_executor.layers.quantization.online.mxfp8 import (
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
QuantKey,
|
||||
kFp8Static128BlockSym,
|
||||
kFp8StaticChannelSym,
|
||||
kFp8StaticTensorSym,
|
||||
kInt8StaticChannelSym,
|
||||
kMxfp8Dynamic,
|
||||
@@ -55,12 +58,14 @@ logger = init_logger(__name__)
|
||||
_ONLINE_LINEAR_METHODS: dict[QuantKey, type] = {
|
||||
kFp8StaticTensorSym: Fp8PerTensorOnlineLinearMethod,
|
||||
kFp8Static128BlockSym: Fp8PerBlockOnlineLinearMethod,
|
||||
kFp8StaticChannelSym: Fp8PtpcOnlineLinearMethod,
|
||||
kMxfp8Dynamic: Mxfp8OnlineLinearMethod,
|
||||
}
|
||||
|
||||
_ONLINE_MOE_METHODS: dict[QuantKey, type] = {
|
||||
kFp8StaticTensorSym: Fp8PerTensorOnlineMoEMethod,
|
||||
kFp8Static128BlockSym: Fp8PerBlockOnlineMoEMethod,
|
||||
kFp8StaticChannelSym: Fp8PtpcOnlineMoEMethod,
|
||||
kMxfp8Dynamic: Mxfp8OnlineMoEMethod,
|
||||
kInt8StaticChannelSym: Int8OnlineMoEMethod,
|
||||
}
|
||||
|
||||
@@ -12,6 +12,7 @@ if TYPE_CHECKING:
|
||||
FusedMoEQuantConfig,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.oracle.fp8 import Fp8MoeBackend
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import QuantKey
|
||||
|
||||
import vllm.envs as envs
|
||||
from vllm import _custom_ops as ops
|
||||
@@ -19,6 +20,7 @@ from vllm.config import get_current_vllm_config
|
||||
from vllm.model_executor.kernels.linear import init_fp8_linear_kernel
|
||||
from vllm.model_executor.kernels.linear.scaled_mm import (
|
||||
CutlassFP8ScaledMMLinearKernel,
|
||||
MarlinFP8ScaledMMLinearKernel,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe import RoutedExperts
|
||||
from vllm.model_executor.layers.fused_moe.oracle.fp8 import (
|
||||
@@ -37,6 +39,7 @@ from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
kFp8DynamicTensorSym,
|
||||
kFp8DynamicTokenSym,
|
||||
kFp8Static128BlockSym,
|
||||
kFp8StaticChannelSym,
|
||||
kFp8StaticTensorSym,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.w8a8_utils import (
|
||||
@@ -270,6 +273,89 @@ class Fp8PerBlockOnlineLinearMethod(_Fp8OnlineLinearBase):
|
||||
)
|
||||
|
||||
|
||||
class Fp8PtpcOnlineLinearMethod(_Fp8OnlineLinearBase):
|
||||
"""Online PTPC FP8 linear quantization.
|
||||
|
||||
Per-output-channel weight scale + dynamic per-token activation scale. The
|
||||
layout matches the llmcompressor's FP8_DYNAMIC recipe, so accuracy
|
||||
is comparable but no pre-quantized checkpoint is required.
|
||||
"""
|
||||
|
||||
weight_quant_key = kFp8StaticChannelSym
|
||||
activation_quant_key = kFp8DynamicTokenSym
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int],
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
):
|
||||
super().create_weights(
|
||||
layer,
|
||||
input_size_per_partition,
|
||||
output_partition_sizes,
|
||||
input_size,
|
||||
output_size,
|
||||
params_dtype,
|
||||
**extra_weight_attrs,
|
||||
)
|
||||
|
||||
self.fp8_linear = init_fp8_linear_kernel(
|
||||
activation_quant_key=self.activation_quant_key,
|
||||
weight_quant_key=self.weight_quant_key,
|
||||
weight_shape=layer.weight.shape,
|
||||
input_dtype=self.input_dtype,
|
||||
out_dtype=self.out_dtype,
|
||||
module_name=self.__class__.__name__,
|
||||
)
|
||||
# PTPC requires per-token activation FP8; MarlinFP8 is W8A16 and
|
||||
# would silently produce a weight-only fp8 model.
|
||||
if isinstance(self.fp8_linear, MarlinFP8ScaledMMLinearKernel):
|
||||
raise ValueError(
|
||||
"FP8 PTPC online quant requires a kernel that honors "
|
||||
"per-token activation quantization; MarlinFP8 is W8A16 "
|
||||
"weight-only. Requires SM89+ for Cutlass FP8 or ROCm MI3xx "
|
||||
"for rowwise scaled_mm."
|
||||
)
|
||||
|
||||
def process_weights_after_loading(self, layer: Module) -> None:
|
||||
if getattr(layer, "_already_called_process_weights_after_loading", False):
|
||||
return
|
||||
|
||||
layer.input_scale = None
|
||||
qweight, weight_scale = ops.scaled_fp8_quant(
|
||||
layer.weight, scale=None, use_per_token_if_dynamic=True
|
||||
)
|
||||
|
||||
replace_parameter(layer, "weight", qweight.t())
|
||||
replace_parameter(layer, "weight_scale", weight_scale)
|
||||
|
||||
self.fp8_linear.process_weights_after_loading(layer)
|
||||
|
||||
layer._already_called_process_weights_after_loading = True
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
# if batch invariant mode is enabled dequant
|
||||
if envs.VLLM_BATCH_INVARIANT and not isinstance(
|
||||
self.fp8_linear, CutlassFP8ScaledMMLinearKernel
|
||||
):
|
||||
weight_dequant = (
|
||||
layer.weight.to(x.dtype) * layer.weight_scale.to(x.dtype).t()
|
||||
)
|
||||
return torch.nn.functional.linear(x, weight_dequant.t(), bias)
|
||||
|
||||
return self.fp8_linear.apply_weights(layer, x, bias)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Online FP8 MoE Methods
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -284,12 +370,17 @@ class _Fp8OnlineMoEBase(OnlineMoEMethodBase):
|
||||
experts_cls: "type[mk.FusedMoEExperts] | None"
|
||||
weight_scale_name: str
|
||||
weight_block_size: list[int] | None
|
||||
per_act_token_quant: bool = False
|
||||
per_out_ch_quant: bool = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
weight_block_size: list[int] | None,
|
||||
layer: torch.nn.Module,
|
||||
weight_key: "QuantKey | None" = None,
|
||||
activation_key: "QuantKey | None" = None,
|
||||
allow_vllm_cutlass: bool = False,
|
||||
):
|
||||
super().__init__(layer.moe_config)
|
||||
self.weight_block_size = weight_block_size
|
||||
@@ -298,20 +389,22 @@ class _Fp8OnlineMoEBase(OnlineMoEMethodBase):
|
||||
"weight_scale_inv" if self.block_quant else "weight_scale"
|
||||
)
|
||||
|
||||
# Set weight key and activation key for kernel compatibility
|
||||
if self.block_quant:
|
||||
weight_key = kFp8Static128BlockSym
|
||||
activation_key = kFp8Dynamic128Sym
|
||||
else:
|
||||
weight_key = kFp8StaticTensorSym
|
||||
activation_key = kFp8DynamicTensorSym
|
||||
# Subclasses may pass explicit kernel keys (PTPC needs channelwise +
|
||||
# per-token).
|
||||
if weight_key is None or activation_key is None:
|
||||
if self.block_quant:
|
||||
weight_key = kFp8Static128BlockSym
|
||||
activation_key = kFp8Dynamic128Sym
|
||||
else:
|
||||
weight_key = kFp8StaticTensorSym
|
||||
activation_key = kFp8DynamicTensorSym
|
||||
|
||||
# Select Fp8 MoE backend
|
||||
self.fp8_backend, self.experts_cls = select_fp8_moe_backend(
|
||||
config=self.moe,
|
||||
weight_key=weight_key,
|
||||
activation_key=activation_key,
|
||||
allow_vllm_cutlass=False,
|
||||
allow_vllm_cutlass=allow_vllm_cutlass,
|
||||
)
|
||||
|
||||
def _setup_kernel(
|
||||
@@ -380,6 +473,8 @@ class _Fp8OnlineMoEBase(OnlineMoEMethodBase):
|
||||
w1_bias=getattr(layer, "w13_bias", None),
|
||||
w2_bias=getattr(layer, "w2_bias", None),
|
||||
block_shape=self.weight_block_size,
|
||||
per_act_token_quant=self.per_act_token_quant,
|
||||
per_out_ch_quant=self.per_out_ch_quant,
|
||||
swiglu_limit=getattr(layer, "swiglu_limit", None),
|
||||
)
|
||||
|
||||
@@ -511,3 +606,89 @@ class Fp8PerBlockOnlineMoEMethod(_Fp8OnlineMoEBase):
|
||||
|
||||
# Prevent duplicate processing (e.g., during weight reload)
|
||||
layer._already_called_process_weights_after_loading = True
|
||||
|
||||
|
||||
class Fp8PtpcOnlineMoEMethod(_Fp8OnlineMoEBase):
|
||||
"""Online PTPC FP8 MoE quantization.
|
||||
|
||||
Quantizes each expert's weights per output channel during loading.
|
||||
Activations are quantized dynamically per token at runtime.
|
||||
"""
|
||||
|
||||
per_act_token_quant: bool = True
|
||||
per_out_ch_quant: bool = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
layer: torch.nn.Module,
|
||||
):
|
||||
from vllm.model_executor.layers.fused_moe.oracle.fp8 import Fp8MoeBackend
|
||||
|
||||
super().__init__(
|
||||
weight_block_size=None,
|
||||
layer=layer,
|
||||
weight_key=kFp8StaticChannelSym,
|
||||
activation_key=kFp8DynamicTokenSym,
|
||||
allow_vllm_cutlass=True,
|
||||
)
|
||||
# Reject backends whose make_fp8_moe_quant_config branch silently
|
||||
# drops per_act_token_quant / per_out_ch_quant or collapses scales:
|
||||
# MARLIN / CPU route through fp8_w8a16_moe_quant_config; FLASHINFER_*
|
||||
# fold scales into a per-tensor alpha (oracle/fp8.py).
|
||||
if self.fp8_backend in (
|
||||
Fp8MoeBackend.MARLIN,
|
||||
Fp8MoeBackend.CPU,
|
||||
Fp8MoeBackend.FLASHINFER_CUTLASS,
|
||||
Fp8MoeBackend.FLASHINFER_TRTLLM,
|
||||
):
|
||||
raise ValueError(
|
||||
f"FP8 PTPC online MoE quant is not supported with the "
|
||||
f"{self.fp8_backend.value} backend, which does not implement "
|
||||
"per-output-channel weight scales."
|
||||
)
|
||||
|
||||
def process_weights_after_loading(self, layer: Module) -> None:
|
||||
if getattr(layer, "_already_called_process_weights_after_loading", False):
|
||||
return
|
||||
|
||||
fp8_dtype = current_platform.fp8_dtype()
|
||||
w13 = torch.empty_like(layer.w13_weight, dtype=fp8_dtype)
|
||||
w2 = torch.empty_like(layer.w2_weight, dtype=fp8_dtype)
|
||||
# Scale's leading dim is taken from the fp8 weight tensor by
|
||||
# construction, so it cannot drift from the weight's expert count
|
||||
# under EP / padded MoE.
|
||||
n_w13 = layer.w13_weight.shape[1]
|
||||
n_w2 = layer.w2_weight.shape[1]
|
||||
w13_scale = torch.ones(
|
||||
w13.shape[0], n_w13, 1, device=w13.device, dtype=torch.float32
|
||||
)
|
||||
w2_scale = torch.ones(
|
||||
w2.shape[0], n_w2, 1, device=w2.device, dtype=torch.float32
|
||||
)
|
||||
layer.w13_input_scale = None
|
||||
layer.w2_input_scale = None
|
||||
|
||||
for expert in range(layer.local_num_experts):
|
||||
w13[expert], w13_scale[expert] = ops.scaled_fp8_quant(
|
||||
layer.w13_weight[expert],
|
||||
scale=None,
|
||||
use_per_token_if_dynamic=True,
|
||||
)
|
||||
w2[expert], w2_scale[expert] = ops.scaled_fp8_quant(
|
||||
layer.w2_weight[expert],
|
||||
scale=None,
|
||||
use_per_token_if_dynamic=True,
|
||||
)
|
||||
|
||||
self._setup_kernel(
|
||||
layer,
|
||||
w13,
|
||||
w2,
|
||||
w13_scale,
|
||||
w2_scale,
|
||||
w13_input_scale=None,
|
||||
w2_input_scale=None,
|
||||
)
|
||||
|
||||
layer._already_called_process_weights_after_loading = True
|
||||
|
||||
Reference in New Issue
Block a user