fix: llama 4 + trtllm gen + fp8 kv cache incompatibility (#12347)
This commit is contained in:
@@ -971,6 +971,13 @@ class ServerArgs:
|
|||||||
logger.warning(
|
logger.warning(
|
||||||
"Use trtllm_mha as attention backend on sm100 for Llama4 model"
|
"Use trtllm_mha as attention backend on sm100 for Llama4 model"
|
||||||
)
|
)
|
||||||
|
if is_sm100_supported() and self.attention_backend == "trtllm_mha":
|
||||||
|
# TODO(brayden): remove this once TRTLLM MHA kernel for FP8 w/ tileSizeKv=128 is available.
|
||||||
|
# This is a Llama 4 specific issue only.
|
||||||
|
self.kv_cache_dtype = "bfloat16"
|
||||||
|
logger.warning(
|
||||||
|
"Setting kv_cache_dtype to bfloat16 for Llama4 with trtllm_mha backend, due to a missing FlashInfer TRTLLM MHA kernel for FP8 KV Cache"
|
||||||
|
)
|
||||||
if is_sm100_supported() and self.moe_runner_backend == "auto":
|
if is_sm100_supported() and self.moe_runner_backend == "auto":
|
||||||
if self.quantization in {"fp8", "modelopt_fp8"}:
|
if self.quantization in {"fp8", "modelopt_fp8"}:
|
||||||
self.moe_runner_backend = "flashinfer_trtllm"
|
self.moe_runner_backend = "flashinfer_trtllm"
|
||||||
|
|||||||
Reference in New Issue
Block a user