[NVIDIA] Support TF32 matmul to improve MiniMax gate gemm performance (#22744)

This commit is contained in:
Trevor Morris
2026-06-23 14:54:54 -07:00
committed by GitHub
parent c864c8d9c2
commit f74a1722e6
3 changed files with 24 additions and 0 deletions
@@ -533,6 +533,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
if self.device == "cpu":
self.init_threads_binding()
# Set float32 matmul precision
if server_args.enable_tf32_matmul:
torch.set_float32_matmul_precision("high")
# Get available memory before model loading.
# Stored for later use by alloc_memory_pool().
self.pre_model_load_memory = self.init_torch_distributed()
+14
View File
@@ -649,6 +649,10 @@ class ServerArgs:
Optional[str],
"Path to the FlashRL quantization profile. Required when using --load-format flash_rl.",
] = None # For flash_rl load format
enable_tf32_matmul: A[
bool,
"Enable float32 matmuls to use TensorFloat32 precision for better performance (via torch.set_float32_matmul_precision). CUDA only.",
] = False
# -------------------------------------------------------------------------
# Memory and scheduling
@@ -4259,6 +4263,10 @@ class ServerArgs:
logger.info(
"Use flashinfer_trtllm as MoE runner backend on sm100 for Glm4MoeForCausalLM"
)
self.enable_tf32_matmul = True
logger.info(
"Enable TF32 matmul for Glm4MoeForCausalLM model to improve gate gemm performance."
)
elif model_arch in [
"FalconH1ForCausalLM",
@@ -4292,6 +4300,12 @@ class ServerArgs:
elif model_arch in ["ZayaForCausalLM"]:
self._handle_mamba_radix_cache(model_arch=model_arch)
elif model_arch in ["MiniMaxM2ForCausalLM"]:
self.enable_tf32_matmul = True
logger.info(
"Enable TF32 matmul for MiniMaxM2ForCausalLM model to improve gate gemm performance."
)
if (
model_arch in ["Qwen3VLForConditionalGeneration"]
and is_hip()