Use FlashInfer tinygemm for GPT-OSS MoE router on SM90+ (#20755)
Co-authored-by: elvischenv <219235043+elvischenv@users.noreply.github.com>
This commit is contained in:
co-authored by
elvischenv
parent
31c35f1c22
commit
bbe25b2412
@@ -75,10 +75,34 @@ from sglang.srt.models.utils import (
|
|||||||
enable_fused_set_kv_buffer,
|
enable_fused_set_kv_buffer,
|
||||||
)
|
)
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import LazyValue, add_prefix, is_npu, make_layers
|
from sglang.srt.utils import (
|
||||||
|
LazyValue,
|
||||||
|
add_prefix,
|
||||||
|
is_blackwell_supported,
|
||||||
|
is_cuda,
|
||||||
|
is_flashinfer_available,
|
||||||
|
is_npu,
|
||||||
|
is_sm90_supported,
|
||||||
|
make_layers,
|
||||||
|
)
|
||||||
from sglang.srt.utils.custom_op import register_custom_op
|
from sglang.srt.utils.custom_op import register_custom_op
|
||||||
|
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
|
_is_cuda = is_cuda()
|
||||||
|
_is_tinygemm_supported = (
|
||||||
|
_is_cuda
|
||||||
|
and is_flashinfer_available()
|
||||||
|
and (is_sm90_supported() or is_blackwell_supported())
|
||||||
|
)
|
||||||
|
|
||||||
|
if _is_tinygemm_supported:
|
||||||
|
try:
|
||||||
|
from flashinfer.gemm import tinygemm_bf16
|
||||||
|
except ImportError:
|
||||||
|
tinygemm_bf16 = None
|
||||||
|
_is_tinygemm_supported = False
|
||||||
|
else:
|
||||||
|
tinygemm_bf16 = None
|
||||||
|
|
||||||
|
|
||||||
class GptOssConfig(PretrainedConfig):
|
class GptOssConfig(PretrainedConfig):
|
||||||
@@ -97,6 +121,45 @@ def get_attention_sliding_window_size(config):
|
|||||||
return config.sliding_window - 1
|
return config.sliding_window - 1
|
||||||
|
|
||||||
|
|
||||||
|
class TinyGemmLinear(ReplicatedLinear):
|
||||||
|
"""ReplicatedLinear with a FlashInfer tinygemm BF16 fast path."""
|
||||||
|
|
||||||
|
def __init__(self, *args, **kwargs):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
self._use_tinygemm = (
|
||||||
|
_is_tinygemm_supported
|
||||||
|
and not self.skip_bias_add
|
||||||
|
and self.weight.is_contiguous()
|
||||||
|
and self.weight.shape[0] % 16 == 0
|
||||||
|
and self.weight.shape[1] % 64 == 0
|
||||||
|
and self.weight.dtype == torch.bfloat16
|
||||||
|
and (
|
||||||
|
self.bias is None
|
||||||
|
or (
|
||||||
|
self.bias.dtype == torch.bfloat16
|
||||||
|
and self.bias.is_contiguous()
|
||||||
|
and self.bias.shape[0] == self.weight.shape[0]
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||||
|
if (
|
||||||
|
self._use_tinygemm
|
||||||
|
and x.ndim == 2
|
||||||
|
and x.is_cuda
|
||||||
|
and x.shape[0] <= 128
|
||||||
|
and x.is_contiguous()
|
||||||
|
and x.shape[1] == self.weight.shape[1]
|
||||||
|
and x.dtype == torch.bfloat16
|
||||||
|
):
|
||||||
|
out = x.new_empty((x.shape[0], self.output_size))
|
||||||
|
tinygemm_bf16(x, self.weight, out, self.bias)
|
||||||
|
return out, None
|
||||||
|
|
||||||
|
return super().forward(x)
|
||||||
|
|
||||||
|
|
||||||
class GptOssSparseMoeBlock(nn.Module):
|
class GptOssSparseMoeBlock(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -147,7 +210,7 @@ class GptOssSparseMoeBlock(nn.Module):
|
|||||||
**extra_kwargs,
|
**extra_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.router = ReplicatedLinear(
|
self.router = TinyGemmLinear(
|
||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
config.num_local_experts,
|
config.num_local_experts,
|
||||||
bias=True,
|
bias=True,
|
||||||
|
|||||||
Reference in New Issue
Block a user