diff --git a/docs_new/docs/advanced_features/server_arguments.mdx b/docs_new/docs/advanced_features/server_arguments.mdx index ddf20a9c0..7cd7fc70c 100644 --- a/docs_new/docs/advanced_features/server_arguments.mdx +++ b/docs_new/docs/advanced_features/server_arguments.mdx @@ -394,6 +394,12 @@ Please consult the documentation below and [server_args.py](https://github.com/s `False` bool flag (set to enable) + + `--enable-tf32-matmul` + Enable float32 matmuls to use TensorFloat32 precision for better performance (via torch.set_float32_matmul_precision). CUDA only. Automatically enabled for MiniMax-M2 and GLM-4 models. + `False` + bool flag (set to enable) + diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 0bf828dbf..eed4279b8 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -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() diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index e46ea92be..125940710 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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()