mirror of
https://git.datalinker.icu/vllm-project/vllm.git
synced 2026-09-06 11:27:00 +08:00
[Attention] Make CutlassMLA the default backend for SM100 (blackwell) (#21626)
Signed-off-by: Alexander Matveev <amatveev@redhat.com> Signed-off-by: mgoin <mgoin64@gmail.com> Co-authored-by: mgoin <mgoin64@gmail.com>
This commit is contained in:
parent
a9b2a1d704
commit
8f605ee309
@ -150,11 +150,28 @@ class CudaPlatformBase(Platform):
|
|||||||
# TODO(lucas): handle this more gracefully
|
# TODO(lucas): handle this more gracefully
|
||||||
# Note: model_config may be None during testing
|
# Note: model_config may be None during testing
|
||||||
if model_config is not None and model_config.use_mla:
|
if model_config is not None and model_config.use_mla:
|
||||||
# if `VLLM_ATTENTION_BACKEND` is not set and we are using MLA, then
|
# If `VLLM_ATTENTION_BACKEND` is not set and we are using MLA,
|
||||||
# we default to FlashMLA backend, so we need to force the blocksize
|
# then we default to FlashMLA backend for non-blackwell GPUs,
|
||||||
# here
|
# else we default to CutlassMLA. For each case, we force the
|
||||||
use_flashmla = (envs.VLLM_ATTENTION_BACKEND is None \
|
# required block_size.
|
||||||
or envs.VLLM_ATTENTION_BACKEND == "FLASHMLA")
|
use_flashmla = False
|
||||||
|
use_cutlass_mla = False
|
||||||
|
|
||||||
|
if envs.VLLM_ATTENTION_BACKEND is None:
|
||||||
|
# Default case
|
||||||
|
if cls.is_device_capability(100):
|
||||||
|
# Blackwell => Force CutlassMLA.
|
||||||
|
use_cutlass_mla = True
|
||||||
|
envs.VLLM_ATTENTION_BACKEND = "CUTLASS_MLA_VLLM_V1"
|
||||||
|
else:
|
||||||
|
# Not Blackwell
|
||||||
|
use_flashmla = True
|
||||||
|
else:
|
||||||
|
# Forced case
|
||||||
|
use_flashmla = (envs.VLLM_ATTENTION_BACKEND == "FLASHMLA")
|
||||||
|
use_cutlass_mla = (
|
||||||
|
envs.VLLM_ATTENTION_BACKEND == "CUTLASS_MLA_VLLM_V1")
|
||||||
|
|
||||||
from vllm.attention.ops.flashmla import is_flashmla_supported
|
from vllm.attention.ops.flashmla import is_flashmla_supported
|
||||||
if use_flashmla and is_flashmla_supported()[0] \
|
if use_flashmla and is_flashmla_supported()[0] \
|
||||||
and cache_config.block_size != 64:
|
and cache_config.block_size != 64:
|
||||||
@ -162,8 +179,6 @@ class CudaPlatformBase(Platform):
|
|||||||
logger.info(
|
logger.info(
|
||||||
"Forcing kv cache block size to 64 for FlashMLA backend.")
|
"Forcing kv cache block size to 64 for FlashMLA backend.")
|
||||||
|
|
||||||
use_cutlass_mla = (envs.VLLM_ATTENTION_BACKEND is not None \
|
|
||||||
and envs.VLLM_ATTENTION_BACKEND == "CUTLASS_MLA_VLLM_V1")
|
|
||||||
if use_cutlass_mla and cache_config.block_size != 128:
|
if use_cutlass_mla and cache_config.block_size != 128:
|
||||||
cache_config.block_size = 128
|
cache_config.block_size = 128
|
||||||
logger.info("Forcing kv cache block size to 128 for "
|
logger.info("Forcing kv cache block size to 128 for "
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user