forked from Karylab-cklius/vllm
Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
cbd94b205f | ||
|
|
007219cffd | ||
|
|
fd97715675 |
@@ -20,8 +20,8 @@ import matplotlib.pyplot as plt
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from vllm.model_executor.layers.fused_moe.batched_deep_gemm_moe import (
|
from vllm.model_executor.layers.fused_moe.batched_masked_silu_mul_quant import (
|
||||||
persistent_masked_m_silu_mul_quant,
|
silu_mul_fp8_quant,
|
||||||
)
|
)
|
||||||
from vllm.triton_utils import tl, triton
|
from vllm.triton_utils import tl, triton
|
||||||
from vllm.utils.deep_gemm import is_deep_gemm_e8m0_used
|
from vllm.utils.deep_gemm import is_deep_gemm_e8m0_used
|
||||||
@@ -500,7 +500,7 @@ for id, strategy in enumerate(strategies):
|
|||||||
|
|
||||||
# SiLU V2 (CUDA kernel) results
|
# SiLU V2 (CUDA kernel) results
|
||||||
time_ms_silu_v2, gflops, gbps, perc = benchmark(
|
time_ms_silu_v2, gflops, gbps, perc = benchmark(
|
||||||
persistent_masked_m_silu_mul_quant,
|
silu_mul_fp8_quant,
|
||||||
E,
|
E,
|
||||||
T,
|
T,
|
||||||
H,
|
H,
|
||||||
|
|||||||
@@ -7,249 +7,28 @@ import torch
|
|||||||
import vllm.model_executor.layers.fused_moe.modular_kernel as mk
|
import vllm.model_executor.layers.fused_moe.modular_kernel as mk
|
||||||
from vllm.forward_context import get_forward_context, is_forward_context_available
|
from vllm.forward_context import get_forward_context, is_forward_context_available
|
||||||
from vllm.logger import init_logger
|
from vllm.logger import init_logger
|
||||||
|
|
||||||
|
# Import the fused silu+mul+fp8_quant kernel for batched masked format
|
||||||
|
from vllm.model_executor.layers.fused_moe.batched_masked_silu_mul_quant import (
|
||||||
|
silu_mul_fp8_quant,
|
||||||
|
)
|
||||||
from vllm.model_executor.layers.fused_moe.config import FusedMoEQuantConfig
|
from vllm.model_executor.layers.fused_moe.config import FusedMoEQuantConfig
|
||||||
from vllm.model_executor.layers.fused_moe.topk_weight_and_reduce import (
|
from vllm.model_executor.layers.fused_moe.topk_weight_and_reduce import (
|
||||||
TopKWeightAndReduceDelegate,
|
TopKWeightAndReduceDelegate,
|
||||||
)
|
)
|
||||||
from vllm.model_executor.layers.fused_moe.utils import _resize_cache
|
from vllm.model_executor.layers.fused_moe.utils import _resize_cache
|
||||||
from vllm.platforms import current_platform
|
from vllm.platforms import current_platform
|
||||||
from vllm.triton_utils import tl, triton
|
|
||||||
from vllm.utils.deep_gemm import (
|
from vllm.utils.deep_gemm import (
|
||||||
DeepGemmQuantScaleFMT,
|
DeepGemmQuantScaleFMT,
|
||||||
fp8_m_grouped_gemm_nt_masked,
|
fp8_m_grouped_gemm_nt_masked,
|
||||||
get_mk_alignment_for_contiguous_layout,
|
get_mk_alignment_for_contiguous_layout,
|
||||||
is_deep_gemm_e8m0_used,
|
is_deep_gemm_e8m0_used,
|
||||||
)
|
)
|
||||||
from vllm.utils.math_utils import cdiv, round_up
|
from vllm.utils.math_utils import round_up
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def scales_shape_stride_dtype(
|
|
||||||
E: int, T: int, G: int, quant_scale_fmt: DeepGemmQuantScaleFMT
|
|
||||||
) -> tuple[tuple[int, ...], tuple[int, ...], torch.dtype]:
|
|
||||||
shape = (E, T, G)
|
|
||||||
strides = (T * G, 1, T)
|
|
||||||
if quant_scale_fmt in [
|
|
||||||
DeepGemmQuantScaleFMT.FLOAT32,
|
|
||||||
DeepGemmQuantScaleFMT.FLOAT32_CEIL_UE8M0,
|
|
||||||
]:
|
|
||||||
return shape, strides, torch.float32
|
|
||||||
|
|
||||||
assert quant_scale_fmt == DeepGemmQuantScaleFMT.UE8M0
|
|
||||||
shape = (E, T, cdiv(G, 4))
|
|
||||||
strides = (T * cdiv(G, 4), 1, T)
|
|
||||||
return shape, strides, torch.int32
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def _silu_mul_fp8_quant_deep_gemm(
|
|
||||||
# Pointers ------------------------------------------------------------
|
|
||||||
input_ptr, # 16-bit activations (E, T, 2*H)
|
|
||||||
y_q_ptr, # fp8 quantized activations (E, T, H)
|
|
||||||
y_s_ptr, # 16-bit scales (E, T, G)
|
|
||||||
counts_ptr, # int32 num tokens per expert (E)
|
|
||||||
# Sizes ---------------------------------------------------------------
|
|
||||||
H: tl.constexpr, # hidden dimension (per output)
|
|
||||||
GROUP_SIZE: tl.constexpr, # elements per group (usually 128)
|
|
||||||
# Strides for input (elements) ---------------------------------------
|
|
||||||
stride_i_e,
|
|
||||||
stride_i_t,
|
|
||||||
stride_i_h,
|
|
||||||
# Strides for y_q (elements) -----------------------------------------
|
|
||||||
stride_yq_e,
|
|
||||||
stride_yq_t,
|
|
||||||
stride_yq_h,
|
|
||||||
# Strides for y_s (elements) -----------------------------------------
|
|
||||||
stride_ys_e,
|
|
||||||
stride_ys_t,
|
|
||||||
stride_ys_g,
|
|
||||||
# Stride for counts (elements)
|
|
||||||
stride_counts_e,
|
|
||||||
# Numeric params ------------------------------------------------------
|
|
||||||
eps: tl.constexpr,
|
|
||||||
fp8_min: tl.constexpr,
|
|
||||||
fp8_max: tl.constexpr,
|
|
||||||
ceil_ue8m0: tl.constexpr,
|
|
||||||
# Meta ---------------------------------------------------------------
|
|
||||||
BLOCK: tl.constexpr,
|
|
||||||
NUM_STAGES: tl.constexpr,
|
|
||||||
):
|
|
||||||
G = H // GROUP_SIZE
|
|
||||||
|
|
||||||
# map program id -> (e, g)
|
|
||||||
pid = tl.program_id(0)
|
|
||||||
e = pid // G
|
|
||||||
g = pid % G
|
|
||||||
|
|
||||||
e = e.to(tl.int64)
|
|
||||||
g = g.to(tl.int64)
|
|
||||||
|
|
||||||
# number of valid tokens for this expert
|
|
||||||
n_tokens = tl.load(counts_ptr + e * stride_counts_e).to(tl.int64)
|
|
||||||
|
|
||||||
cols = tl.arange(0, BLOCK).to(tl.int64)
|
|
||||||
mask = cols < BLOCK
|
|
||||||
|
|
||||||
base_input_offset = e * stride_i_e + g * GROUP_SIZE * stride_i_h
|
|
||||||
base_gate_offset = base_input_offset + cols * stride_i_h
|
|
||||||
base_up_offset = base_input_offset + H * stride_i_h + cols * stride_i_h
|
|
||||||
base_yq_offset = e * stride_yq_e + g * GROUP_SIZE * stride_yq_h + cols * stride_yq_h
|
|
||||||
base_ys_offset = e * stride_ys_e + g * stride_ys_g
|
|
||||||
|
|
||||||
for t in tl.range(0, n_tokens, num_stages=NUM_STAGES):
|
|
||||||
gate = tl.load(
|
|
||||||
input_ptr + base_gate_offset + t * stride_i_t, mask=mask, other=0.0
|
|
||||||
).to(tl.float32)
|
|
||||||
up = tl.load(input_ptr + base_up_offset + t * stride_i_t, mask=mask, other=0.0)
|
|
||||||
|
|
||||||
gate = gate * (1.0 / (1.0 + tl.exp(-gate)))
|
|
||||||
y = gate * up
|
|
||||||
|
|
||||||
y_s = tl.maximum(tl.max(tl.abs(y)), eps) / fp8_max
|
|
||||||
if ceil_ue8m0:
|
|
||||||
y_s = tl.exp2(tl.ceil(tl.log2(y_s)))
|
|
||||||
|
|
||||||
y_q = tl.clamp(y / y_s, fp8_min, fp8_max).to(y_q_ptr.dtype.element_ty)
|
|
||||||
|
|
||||||
tl.store(y_q_ptr + base_yq_offset + t * stride_yq_t, y_q, mask=mask)
|
|
||||||
tl.store(y_s_ptr + base_ys_offset + t * stride_ys_t, y_s)
|
|
||||||
|
|
||||||
|
|
||||||
def persistent_masked_m_silu_mul_quant(
|
|
||||||
y: torch.Tensor, # (E, T, 2*H)
|
|
||||||
tokens_per_expert: torch.Tensor, # (E,) number of valid tokens per expert
|
|
||||||
num_parallel_tokens=16,
|
|
||||||
group_size: int = 128,
|
|
||||||
quant_scale_fmt: DeepGemmQuantScaleFMT = DeepGemmQuantScaleFMT.FLOAT32,
|
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
||||||
"""Quantize silu(y[..., :H]) * y[..., H:] to FP8 with group per-token scales
|
|
||||||
y has shape (E, T, 2*H). The first half of the last dimension is
|
|
||||||
silu-activated, multiplied by the second half, then quantized into FP8.
|
|
||||||
We launch a fixed grid of threads to accommodate CUDA graphs. Let `P2`
|
|
||||||
be a parallelization factor for persistent_masked_m_silu_mul_quant over the
|
|
||||||
hidden dimension.
|
|
||||||
|
|
||||||
Let `expert_offsets = [0] + [num_tokens.cumsum()]` and
|
|
||||||
`total_tokens = expert_offsets[-1]`.
|
|
||||||
persistent_masked_m_silu_mul_quant launches `total_tokens x P2` number of
|
|
||||||
thread blocks. Each thread block contains `NUM_WARPS` warps.
|
|
||||||
|
|
||||||
Every thread block needs to find it's corresponding expert by warp-parallel scanning
|
|
||||||
over the `expert_offsets` array.
|
|
||||||
|
|
||||||
The i-th warp in the first thread block processes
|
|
||||||
`[i * warp_chunk_size, (i + 1) * warp_chunk_size]` groups
|
|
||||||
sequentially, where `warp_chunk_size = ((H / GROUP_SIZE) / P2) / NUM_WARPS`,
|
|
||||||
pipelining loads and computes.
|
|
||||||
|
|
||||||
The shared memory layout for 4 warps with a 2-stage pipeline for SiLU V2
|
|
||||||
can is visualized like so:
|
|
||||||
|
|
||||||
stage0 stage1
|
|
||||||
┌─────┬───┬─────┬───┬─────┬───┬─────┬───┬─────┬───┬─────┬───┬─────┬───┬─────┬───┐
|
|
||||||
│gate0│up0│gate1│up1│gate2│up2│gate3│up3│gate0│up0│gate1│up1│gate2│up2│gate3│up3│
|
|
||||||
└─────┴───┴─────┴───┴─────┴───┴─────┴───┴─────┴───┴─────┴───┴─────┴───┴─────┴───┘
|
|
||||||
|
|
||||||
with the main difference between V1 and V2 being the global load
|
|
||||||
stride between warps, and between half-warps. Regarding the latter stride,
|
|
||||||
we assign the first half warp of every warp for `gate` loads and the second
|
|
||||||
half-warp to `up` loads.
|
|
||||||
|
|
||||||
Returns `(y_q, y_s)` where
|
|
||||||
* `y_q`: FP8 tensor, shape (E, T, H), same layout as y[..., :H]
|
|
||||||
* `y_s` depends on quant_scale_fmt,
|
|
||||||
- quant_scale_fmt == FLOAT32,
|
|
||||||
`y_s`: FP32 tensor, shape (E, T, H // group_size), strides (T*G, 1, T)
|
|
||||||
- quant_scale_fmt == E8M0,
|
|
||||||
`y_s`: Int32 tensor, shape (E, T, H // group_size // 4), strides (T*G, 1, T)
|
|
||||||
- quant_scale_fmt == E8M0_FLOAT32_SPARSE
|
|
||||||
`y_s`: FP32 tensor, shape (E, T, H // group_size), strides (T*G, 1, T)
|
|
||||||
Let NUM_WARPS be the number of warps in a single thread block and
|
|
||||||
`GROUP_SIZE = 128` be the size of the quantization group.
|
|
||||||
"""
|
|
||||||
assert y.ndim == 3, "y must be (E, T, 2*H)"
|
|
||||||
E, T, H2 = y.shape
|
|
||||||
assert H2 % 2 == 0, "last dim of y must be even (2*H)"
|
|
||||||
H = H2 // 2
|
|
||||||
G = (H + group_size - 1) // group_size
|
|
||||||
assert H % 8 == 0, "H must be divisible by 8"
|
|
||||||
assert group_size == 128, "H must be divisible by 8"
|
|
||||||
assert tokens_per_expert.ndim == 1 and tokens_per_expert.shape[0] == E
|
|
||||||
|
|
||||||
tokens_per_expert = tokens_per_expert.to(device=y.device, dtype=torch.int32)
|
|
||||||
|
|
||||||
fp8_dtype = torch.float8_e4m3fn
|
|
||||||
y_q = torch.empty((E, T, H), dtype=fp8_dtype, device=y.device)
|
|
||||||
|
|
||||||
ys_shape, ys_strides, ys_dtype = scales_shape_stride_dtype(E, T, G, quant_scale_fmt)
|
|
||||||
y_s = torch.empty_strided(
|
|
||||||
ys_shape,
|
|
||||||
ys_strides,
|
|
||||||
dtype=ys_dtype,
|
|
||||||
device=y.device,
|
|
||||||
)
|
|
||||||
|
|
||||||
ceil_ue8m0 = quant_scale_fmt in [
|
|
||||||
DeepGemmQuantScaleFMT.FLOAT32_CEIL_UE8M0,
|
|
||||||
DeepGemmQuantScaleFMT.UE8M0,
|
|
||||||
]
|
|
||||||
|
|
||||||
cuda_arch = current_platform.get_device_capability(
|
|
||||||
device_id=y.device.index
|
|
||||||
).to_int()
|
|
||||||
|
|
||||||
if cuda_arch >= 80:
|
|
||||||
torch.ops._C.persistent_masked_m_silu_mul_quant(
|
|
||||||
y, tokens_per_expert, y_q, y_s, ceil_ue8m0
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
stride_cnt_e = tokens_per_expert.stride()[0]
|
|
||||||
|
|
||||||
# Static grid over experts and H-groups.
|
|
||||||
# A loop inside the kernel handles the token dim
|
|
||||||
grid = (E * G,)
|
|
||||||
# strides (elements)
|
|
||||||
stride_i_e, stride_i_t, stride_i_h = y.stride()
|
|
||||||
stride_yq_e, stride_yq_t, stride_yq_h = y_q.stride()
|
|
||||||
|
|
||||||
f_info = torch.finfo(fp8_dtype)
|
|
||||||
fp8_max = f_info.max
|
|
||||||
fp8_min = f_info.min
|
|
||||||
eps: float = 1e-10
|
|
||||||
assert y_s.dtype == torch.float32, (
|
|
||||||
"_silu_mul_fp8_quant_deep_gemm does"
|
|
||||||
"not support {y_s.dtype} scales. Only torch.float32 supported."
|
|
||||||
)
|
|
||||||
_silu_mul_fp8_quant_deep_gemm[grid](
|
|
||||||
y,
|
|
||||||
y_q,
|
|
||||||
y_s,
|
|
||||||
tokens_per_expert,
|
|
||||||
H,
|
|
||||||
group_size,
|
|
||||||
stride_i_e,
|
|
||||||
stride_i_t,
|
|
||||||
stride_i_h,
|
|
||||||
stride_yq_e,
|
|
||||||
stride_yq_t,
|
|
||||||
stride_yq_h,
|
|
||||||
ys_strides[0],
|
|
||||||
ys_strides[1],
|
|
||||||
ys_strides[2],
|
|
||||||
stride_cnt_e,
|
|
||||||
eps,
|
|
||||||
fp8_min,
|
|
||||||
fp8_max,
|
|
||||||
ceil_ue8m0,
|
|
||||||
BLOCK=group_size,
|
|
||||||
NUM_STAGES=4,
|
|
||||||
num_warps=1,
|
|
||||||
)
|
|
||||||
|
|
||||||
return y_q, y_s
|
|
||||||
|
|
||||||
|
|
||||||
class BatchedDeepGemmExperts(mk.FusedMoEPermuteExpertsUnpermute):
|
class BatchedDeepGemmExperts(mk.FusedMoEPermuteExpertsUnpermute):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -396,7 +175,7 @@ class BatchedDeepGemmExperts(mk.FusedMoEPermuteExpertsUnpermute):
|
|||||||
)
|
)
|
||||||
|
|
||||||
quant_scale_fmt = DeepGemmQuantScaleFMT.from_oracle()
|
quant_scale_fmt = DeepGemmQuantScaleFMT.from_oracle()
|
||||||
a2q, a2q_scale = persistent_masked_m_silu_mul_quant(
|
a2q, a2q_scale = silu_mul_fp8_quant(
|
||||||
workspace1,
|
workspace1,
|
||||||
expert_num_tokens,
|
expert_num_tokens,
|
||||||
quant_scale_fmt=quant_scale_fmt,
|
quant_scale_fmt=quant_scale_fmt,
|
||||||
|
|||||||
@@ -0,0 +1,251 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||||
|
"""Fused SiLU + Mul + FP8 Quantization kernel.
|
||||||
|
|
||||||
|
This module provides a fused kernel that combines:
|
||||||
|
1. SiLU activation on the first half of the hidden dimension
|
||||||
|
2. Element-wise multiplication with the second half (gated)
|
||||||
|
3. FP8 quantization with per-group scales
|
||||||
|
|
||||||
|
The kernel is used by both BatchedDeepGemmExperts and BatchedTritonExperts
|
||||||
|
for efficient activation and quantization in MoE layers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from vllm.platforms import current_platform
|
||||||
|
from vllm.triton_utils import tl, triton
|
||||||
|
from vllm.utils.deep_gemm import DeepGemmQuantScaleFMT
|
||||||
|
from vllm.utils.math_utils import cdiv
|
||||||
|
|
||||||
|
|
||||||
|
def scales_shape_stride_dtype(
|
||||||
|
E: int, T: int, G: int, quant_scale_fmt: DeepGemmQuantScaleFMT
|
||||||
|
) -> tuple[tuple[int, ...], tuple[int, ...], torch.dtype]:
|
||||||
|
"""Compute shape, strides, and dtype for quantization scales tensor.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
E: Number of experts
|
||||||
|
T: Max tokens per expert
|
||||||
|
G: Number of groups (H // group_size)
|
||||||
|
quant_scale_fmt: Scale format (FLOAT32, FLOAT32_CEIL_UE8M0, or UE8M0)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (shape, strides, dtype) for the scales tensor
|
||||||
|
"""
|
||||||
|
shape = (E, T, G)
|
||||||
|
strides = (T * G, 1, T)
|
||||||
|
if quant_scale_fmt in [
|
||||||
|
DeepGemmQuantScaleFMT.FLOAT32,
|
||||||
|
DeepGemmQuantScaleFMT.FLOAT32_CEIL_UE8M0,
|
||||||
|
]:
|
||||||
|
return shape, strides, torch.float32
|
||||||
|
|
||||||
|
assert quant_scale_fmt == DeepGemmQuantScaleFMT.UE8M0
|
||||||
|
shape = (E, T, cdiv(G, 4))
|
||||||
|
strides = (T * cdiv(G, 4), 1, T)
|
||||||
|
return shape, strides, torch.int32
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _silu_mul_fp8_quant_kernel(
|
||||||
|
# Pointers ------------------------------------------------------------
|
||||||
|
input_ptr, # 16-bit activations (E, T, 2*H)
|
||||||
|
y_q_ptr, # fp8 quantized activations (E, T, H)
|
||||||
|
y_s_ptr, # 16-bit scales (E, T, G)
|
||||||
|
counts_ptr, # int32 num tokens per expert (E)
|
||||||
|
# Sizes ---------------------------------------------------------------
|
||||||
|
H: tl.constexpr, # hidden dimension (per output)
|
||||||
|
GROUP_SIZE: tl.constexpr, # elements per group (usually 128)
|
||||||
|
# Strides for input (elements) ---------------------------------------
|
||||||
|
stride_i_e,
|
||||||
|
stride_i_t,
|
||||||
|
stride_i_h,
|
||||||
|
# Strides for y_q (elements) -----------------------------------------
|
||||||
|
stride_yq_e,
|
||||||
|
stride_yq_t,
|
||||||
|
stride_yq_h,
|
||||||
|
# Strides for y_s (elements) -----------------------------------------
|
||||||
|
stride_ys_e,
|
||||||
|
stride_ys_t,
|
||||||
|
stride_ys_g,
|
||||||
|
# Stride for counts (elements)
|
||||||
|
stride_counts_e,
|
||||||
|
# Numeric params ------------------------------------------------------
|
||||||
|
eps: tl.constexpr,
|
||||||
|
fp8_min: tl.constexpr,
|
||||||
|
fp8_max: tl.constexpr,
|
||||||
|
ceil_ue8m0: tl.constexpr,
|
||||||
|
# Meta ---------------------------------------------------------------
|
||||||
|
BLOCK: tl.constexpr,
|
||||||
|
NUM_STAGES: tl.constexpr,
|
||||||
|
):
|
||||||
|
"""Triton kernel for fused SiLU + mul + FP8 quantization.
|
||||||
|
|
||||||
|
For each expert and group, this kernel:
|
||||||
|
1. Loads gate and up projections from input (shape: E, T, 2*H)
|
||||||
|
2. Applies SiLU: gate = gate * sigmoid(gate)
|
||||||
|
3. Applies gated multiplication: y = gate * up
|
||||||
|
4. Quantizes to FP8 with per-group scales
|
||||||
|
"""
|
||||||
|
G = H // GROUP_SIZE
|
||||||
|
|
||||||
|
# map program id -> (e, g)
|
||||||
|
pid = tl.program_id(0)
|
||||||
|
e = pid // G
|
||||||
|
g = pid % G
|
||||||
|
|
||||||
|
e = e.to(tl.int64)
|
||||||
|
g = g.to(tl.int64)
|
||||||
|
|
||||||
|
# number of valid tokens for this expert
|
||||||
|
n_tokens = tl.load(counts_ptr + e * stride_counts_e).to(tl.int64)
|
||||||
|
|
||||||
|
cols = tl.arange(0, BLOCK).to(tl.int64)
|
||||||
|
mask = cols < BLOCK
|
||||||
|
|
||||||
|
base_input_offset = e * stride_i_e + g * GROUP_SIZE * stride_i_h
|
||||||
|
base_gate_offset = base_input_offset + cols * stride_i_h
|
||||||
|
base_up_offset = base_input_offset + H * stride_i_h + cols * stride_i_h
|
||||||
|
base_yq_offset = e * stride_yq_e + g * GROUP_SIZE * stride_yq_h + cols * stride_yq_h
|
||||||
|
base_ys_offset = e * stride_ys_e + g * stride_ys_g
|
||||||
|
|
||||||
|
for t in tl.range(0, n_tokens, num_stages=NUM_STAGES):
|
||||||
|
gate = tl.load(
|
||||||
|
input_ptr + base_gate_offset + t * stride_i_t, mask=mask, other=0.0
|
||||||
|
).to(tl.float32)
|
||||||
|
up = tl.load(input_ptr + base_up_offset + t * stride_i_t, mask=mask, other=0.0)
|
||||||
|
|
||||||
|
# SiLU activation: gate * sigmoid(gate)
|
||||||
|
gate = gate * (1.0 / (1.0 + tl.exp(-gate)))
|
||||||
|
# Gated multiplication
|
||||||
|
y = gate * up
|
||||||
|
|
||||||
|
# Compute per-group scale
|
||||||
|
y_s = tl.maximum(tl.max(tl.abs(y)), eps) / fp8_max
|
||||||
|
if ceil_ue8m0:
|
||||||
|
# Round scale to power of 2 for UE8M0 format
|
||||||
|
y_s = tl.exp2(tl.ceil(tl.log2(y_s)))
|
||||||
|
|
||||||
|
# Quantize to FP8
|
||||||
|
y_q = tl.clamp(y / y_s, fp8_min, fp8_max).to(y_q_ptr.dtype.element_ty)
|
||||||
|
|
||||||
|
tl.store(y_q_ptr + base_yq_offset + t * stride_yq_t, y_q, mask=mask)
|
||||||
|
tl.store(y_s_ptr + base_ys_offset + t * stride_ys_t, y_s)
|
||||||
|
|
||||||
|
|
||||||
|
def silu_mul_fp8_quant(
|
||||||
|
y: torch.Tensor, # (E, T, 2*H)
|
||||||
|
tokens_per_expert: torch.Tensor, # (E,) number of valid tokens per expert
|
||||||
|
group_size: int = 128,
|
||||||
|
quant_scale_fmt: DeepGemmQuantScaleFMT = DeepGemmQuantScaleFMT.FLOAT32,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""Fused SiLU activation + gated multiplication + FP8 quantization.
|
||||||
|
|
||||||
|
Computes: quantize(silu(y[..., :H]) * y[..., H:])
|
||||||
|
|
||||||
|
y has shape (E, T, 2*H). The first half of the last dimension is
|
||||||
|
silu-activated, multiplied by the second half, then quantized into FP8.
|
||||||
|
|
||||||
|
On CUDA SM80+ devices, uses an optimized CUDA kernel.
|
||||||
|
On other platforms (including ROCm), uses a Triton kernel.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
y: Input tensor of shape (E, T, 2*H) where:
|
||||||
|
- E = number of experts
|
||||||
|
- T = max tokens per expert
|
||||||
|
- 2*H = gate and up projections concatenated
|
||||||
|
tokens_per_expert: Number of valid tokens per expert, shape (E,)
|
||||||
|
group_size: Quantization group size (default 128)
|
||||||
|
quant_scale_fmt: Scale format for quantization
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (y_q, y_s) where:
|
||||||
|
- y_q: FP8 tensor, shape (E, T, H)
|
||||||
|
- y_s: Scales tensor, shape and dtype depend on quant_scale_fmt:
|
||||||
|
- FLOAT32: FP32 tensor, shape (E, T, G), strides (T*G, 1, T)
|
||||||
|
- FLOAT32_CEIL_UE8M0: FP32 tensor, shape (E, T, G), strides (T*G, 1, T)
|
||||||
|
- UE8M0: Int32 tensor, shape (E, T, G//4), strides (T*G//4, 1, T)
|
||||||
|
"""
|
||||||
|
assert y.ndim == 3, "y must be (E, T, 2*H)"
|
||||||
|
E, T, H2 = y.shape
|
||||||
|
assert H2 % 2 == 0, "last dim of y must be even (2*H)"
|
||||||
|
H = H2 // 2
|
||||||
|
G = (H + group_size - 1) // group_size
|
||||||
|
assert H % 8 == 0, "H must be divisible by 8"
|
||||||
|
assert group_size == 128, "group_size must be 128"
|
||||||
|
assert tokens_per_expert.ndim == 1 and tokens_per_expert.shape[0] == E
|
||||||
|
|
||||||
|
tokens_per_expert = tokens_per_expert.to(device=y.device, dtype=torch.int32)
|
||||||
|
|
||||||
|
fp8_dtype = torch.float8_e4m3fn
|
||||||
|
y_q = torch.empty((E, T, H), dtype=fp8_dtype, device=y.device)
|
||||||
|
|
||||||
|
ys_shape, ys_strides, ys_dtype = scales_shape_stride_dtype(E, T, G, quant_scale_fmt)
|
||||||
|
y_s = torch.empty_strided(
|
||||||
|
ys_shape,
|
||||||
|
ys_strides,
|
||||||
|
dtype=ys_dtype,
|
||||||
|
device=y.device,
|
||||||
|
)
|
||||||
|
|
||||||
|
ceil_ue8m0 = quant_scale_fmt in [
|
||||||
|
DeepGemmQuantScaleFMT.FLOAT32_CEIL_UE8M0,
|
||||||
|
DeepGemmQuantScaleFMT.UE8M0,
|
||||||
|
]
|
||||||
|
|
||||||
|
cuda_arch = current_platform.get_device_capability(
|
||||||
|
device_id=y.device.index
|
||||||
|
).to_int()
|
||||||
|
|
||||||
|
if cuda_arch >= 80:
|
||||||
|
# Use optimized CUDA kernel on SM80+ devices
|
||||||
|
torch.ops._C.persistent_masked_m_silu_mul_quant(
|
||||||
|
y, tokens_per_expert, y_q, y_s, ceil_ue8m0
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Use Triton kernel on other platforms (including ROCm)
|
||||||
|
stride_cnt_e = tokens_per_expert.stride()[0]
|
||||||
|
|
||||||
|
# Static grid over experts and H-groups.
|
||||||
|
# A loop inside the kernel handles the token dim
|
||||||
|
grid = (E * G,)
|
||||||
|
# strides (elements)
|
||||||
|
stride_i_e, stride_i_t, stride_i_h = y.stride()
|
||||||
|
stride_yq_e, stride_yq_t, stride_yq_h = y_q.stride()
|
||||||
|
|
||||||
|
f_info = torch.finfo(fp8_dtype)
|
||||||
|
fp8_max = f_info.max
|
||||||
|
fp8_min = f_info.min
|
||||||
|
eps: float = 1e-10
|
||||||
|
assert y_s.dtype == torch.float32, (
|
||||||
|
f"_silu_mul_fp8_quant_kernel does not support {y_s.dtype} scales. "
|
||||||
|
"Only torch.float32 supported."
|
||||||
|
)
|
||||||
|
_silu_mul_fp8_quant_kernel[grid](
|
||||||
|
y,
|
||||||
|
y_q,
|
||||||
|
y_s,
|
||||||
|
tokens_per_expert,
|
||||||
|
H,
|
||||||
|
group_size,
|
||||||
|
stride_i_e,
|
||||||
|
stride_i_t,
|
||||||
|
stride_i_h,
|
||||||
|
stride_yq_e,
|
||||||
|
stride_yq_t,
|
||||||
|
stride_yq_h,
|
||||||
|
ys_strides[0],
|
||||||
|
ys_strides[1],
|
||||||
|
ys_strides[2],
|
||||||
|
stride_cnt_e,
|
||||||
|
eps,
|
||||||
|
fp8_min,
|
||||||
|
fp8_max,
|
||||||
|
ceil_ue8m0,
|
||||||
|
BLOCK=group_size,
|
||||||
|
NUM_STAGES=4,
|
||||||
|
num_warps=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
return y_q, y_s
|
||||||
@@ -5,6 +5,12 @@
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
import vllm.model_executor.layers.fused_moe.modular_kernel as mk
|
import vllm.model_executor.layers.fused_moe.modular_kernel as mk
|
||||||
|
|
||||||
|
# Import the fused silu+mul+fp8_quant kernel for batched masked format.
|
||||||
|
# This wrapper calls the Triton kernel on non-SM80+ platforms (including ROCm).
|
||||||
|
from vllm.model_executor.layers.fused_moe.batched_masked_silu_mul_quant import (
|
||||||
|
silu_mul_fp8_quant,
|
||||||
|
)
|
||||||
from vllm.model_executor.layers.fused_moe.config import FusedMoEQuantConfig
|
from vllm.model_executor.layers.fused_moe.config import FusedMoEQuantConfig
|
||||||
from vllm.model_executor.layers.fused_moe.fused_moe import try_get_optimal_moe_config
|
from vllm.model_executor.layers.fused_moe.fused_moe import try_get_optimal_moe_config
|
||||||
from vllm.model_executor.layers.fused_moe.topk_weight_and_reduce import (
|
from vllm.model_executor.layers.fused_moe.topk_weight_and_reduce import (
|
||||||
@@ -19,6 +25,10 @@ from vllm.model_executor.layers.fused_moe.utils import (
|
|||||||
)
|
)
|
||||||
from vllm.model_executor.layers.quantization.utils.quant_utils import group_broadcast
|
from vllm.model_executor.layers.quantization.utils.quant_utils import group_broadcast
|
||||||
from vllm.triton_utils import tl, triton
|
from vllm.triton_utils import tl, triton
|
||||||
|
from vllm.utils.deep_gemm import DeepGemmQuantScaleFMT
|
||||||
|
|
||||||
|
# Default group size for FP8 block quantization
|
||||||
|
FUSED_QUANT_GROUP_SIZE = 128
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
@@ -979,26 +989,48 @@ class BatchedTritonExperts(mk.FusedMoEPermuteExpertsUnpermute):
|
|||||||
block_shape=self.block_shape,
|
block_shape=self.block_shape,
|
||||||
)
|
)
|
||||||
|
|
||||||
intermediate_cache2.fill_(0)
|
# Check if we can use the fused silu + mul + fp8 quant kernel.
|
||||||
|
# N here is the output dim of w1 (= 2*H for gated activations).
|
||||||
# TODO (bnell): use triton utility from batched deep gemm.
|
# The kernel requires H % group_size == 0, so we check N % (2*group_size) == 0.
|
||||||
self.activation(
|
use_fused_silu_mul_quant = (
|
||||||
activation,
|
activation == "silu"
|
||||||
intermediate_cache2.view(-1, activation_out_dim),
|
and self.quant_config.use_fp8_w8a8
|
||||||
intermediate_cache1.view(-1, N),
|
and self.block_shape is not None
|
||||||
|
and self.block_shape[1] == FUSED_QUANT_GROUP_SIZE
|
||||||
|
and not self.per_act_token_quant
|
||||||
|
and N % (2 * FUSED_QUANT_GROUP_SIZE) == 0
|
||||||
)
|
)
|
||||||
|
|
||||||
qintermediate_cache2, a2q_scale = batched_moe_kernel_quantize_input(
|
if use_fused_silu_mul_quant:
|
||||||
intermediate_cache2,
|
# Fused path: silu + mul + fp8 quant in one kernel
|
||||||
a2_scale,
|
# intermediate_cache1 has shape (E, max_num_tokens, N) where N = 2*H
|
||||||
max_num_tokens,
|
qintermediate_cache2, a2q_scale = silu_mul_fp8_quant(
|
||||||
E,
|
intermediate_cache1,
|
||||||
N,
|
expert_num_tokens,
|
||||||
expert_num_tokens,
|
group_size=FUSED_QUANT_GROUP_SIZE,
|
||||||
self.quant_dtype,
|
quant_scale_fmt=DeepGemmQuantScaleFMT.FLOAT32,
|
||||||
self.per_act_token_quant,
|
)
|
||||||
self.block_shape,
|
else:
|
||||||
)
|
# Unfused path: separate activation and quantization
|
||||||
|
intermediate_cache2.fill_(0)
|
||||||
|
|
||||||
|
self.activation(
|
||||||
|
activation,
|
||||||
|
intermediate_cache2.view(-1, activation_out_dim),
|
||||||
|
intermediate_cache1.view(-1, N),
|
||||||
|
)
|
||||||
|
|
||||||
|
qintermediate_cache2, a2q_scale = batched_moe_kernel_quantize_input(
|
||||||
|
intermediate_cache2,
|
||||||
|
a2_scale,
|
||||||
|
max_num_tokens,
|
||||||
|
E,
|
||||||
|
N,
|
||||||
|
expert_num_tokens,
|
||||||
|
self.quant_dtype,
|
||||||
|
self.per_act_token_quant,
|
||||||
|
self.block_shape,
|
||||||
|
)
|
||||||
|
|
||||||
invoke_moe_batched_triton_kernel(
|
invoke_moe_batched_triton_kernel(
|
||||||
A=qintermediate_cache2,
|
A=qintermediate_cache2,
|
||||||
|
|||||||
Reference in New Issue
Block a user