From 9ced8d0981aae63f5bba54206d03b10c26f8eef7 Mon Sep 17 00:00:00 2001 From: ilyasher-harmonic Date: Tue, 11 Aug 2026 15:14:50 -0700 Subject: [PATCH] Optimize FP32 LM head for bf16/fp16 (#32370) --- python/sglang/srt/layers/logits_processor.py | 20 +++++- test/registered/rl/test_fp32_lm_head.py | 71 ++++++++++++++++++-- 2 files changed, 84 insertions(+), 7 deletions(-) diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index be995d078..7cea446d9 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -705,9 +705,25 @@ class LogitsProcessor(nn.Module): elif hasattr(lm_head, "weight"): # Normal linear layer if self.use_fp32_lm_head: - logits = torch.matmul( - hidden_states.to(torch.float32), lm_head.weight.to(torch.float32).T + # Avoid materializing FP32 copies for same-dtype CUDA FP16/BF16 + # inputs. Retain explicit FP32 casts for unsupported devices or + # dtype combinations. + use_mm_out_dtype = ( + hidden_states.is_cuda + and hidden_states.dtype == lm_head.weight.dtype + and hidden_states.dtype in (torch.float16, torch.bfloat16) ) + if use_mm_out_dtype: + logits = torch.mm( + hidden_states, + lm_head.weight.T, + out_dtype=torch.float32, + ) + else: + logits = torch.matmul( + hidden_states.to(torch.float32), + lm_head.weight.to(torch.float32).T, + ) elif use_intel_amx_backend(lm_head): logits = torch.ops.sgl_kernel.weight_packed_linear( hidden_states.to(lm_head.weight.dtype), diff --git a/test/registered/rl/test_fp32_lm_head.py b/test/registered/rl/test_fp32_lm_head.py index de7da53da..685b0b969 100644 --- a/test/registered/rl/test_fp32_lm_head.py +++ b/test/registered/rl/test_fp32_lm_head.py @@ -54,6 +54,7 @@ class TestLMHeadFP32(unittest.TestCase): weights_dtype, expected_a_dtype, expected_b_dtype, + expected_operation, ): device = get_device() BATCH_SIZE, HIDDEN_SIZE, VOCAB_SIZE = 2, 64, 128 @@ -65,6 +66,7 @@ class TestLMHeadFP32(unittest.TestCase): logprocessor = self._make_logprocessor(VOCAB_SIZE, enable_fp32) original_matmul = torch.matmul + original_mm = torch.mm original_linear = F.linear state = { @@ -72,13 +74,31 @@ class TestLMHeadFP32(unittest.TestCase): "operation": None, # Which operation was captured ("matmul" or "linear") "a": None, # The dtype of the first input tensor to the operation "b": None, # The dtype of the second input tensor to the operation + "out_dtype": None, } def probe_matmul(a, b, *args, **kw): if not state["called"]: - state.update(called=True, operation="matmul", a=a.dtype, b=b.dtype) + state.update( + called=True, + operation="matmul", + a=a.dtype, + b=b.dtype, + out_dtype=kw.get("out_dtype"), + ) return original_matmul(a, b, *args, **kw) + def probe_mm(a, b, *args, **kw): + if not state["called"]: + state.update( + called=True, + operation="mm", + a=a.dtype, + b=b.dtype, + out_dtype=kw.get("out_dtype"), + ) + return original_mm(a, b, *args, **kw) + def probe_linear(x, w, bias=None): if not state["called"]: state.update(called=True, ooperationp="linear", a=x.dtype, b=w.dtype) @@ -86,30 +106,71 @@ class TestLMHeadFP32(unittest.TestCase): with ( patch("torch.matmul", new=probe_matmul), + patch("torch.mm", new=probe_mm), patch("torch.nn.functional.linear", new=probe_linear), ): logits = logprocessor._get_logits(hidden_state, head, meta) self.assertEqual(hidden_state.dtype, hidden_state_dtype) self.assertTrue(state["called"], "no call lm head matlmul/linear") + self.assertEqual(state["operation"], expected_operation) self.assertEqual(state["a"], expected_a_dtype) self.assertEqual(state["b"], expected_b_dtype) + self.assertEqual( + state["out_dtype"], + torch.float32 if expected_operation == "mm" else None, + ) def test_flag_true_fp16_activations(self): - self._run_case(torch.float16, True, torch.float16, torch.float32, torch.float32) + expected_operation = "mm" if torch.cuda.is_available() else "matmul" + expected_dtype = ( + torch.float32 if expected_operation == "matmul" else torch.float16 + ) + self._run_case( + torch.float16, + True, + torch.float16, + expected_dtype, + expected_dtype, + expected_operation, + ) def test_flag_true_bf16_activations(self): + expected_operation = "mm" if torch.cuda.is_available() else "matmul" + expected_dtype = ( + torch.float32 if expected_operation == "matmul" else torch.bfloat16 + ) self._run_case( - torch.bfloat16, True, torch.bfloat16, torch.float32, torch.float32 + torch.bfloat16, + True, + torch.bfloat16, + expected_dtype, + expected_dtype, + expected_operation, + ) + + def test_flag_true_fp32_falls_back_to_explicit_fp32_matmul(self): + self._run_case( + torch.float32, + True, + torch.float32, + torch.float32, + torch.float32, + "matmul", ) def test_flag_false_fp16_path(self): self._run_case( - torch.float16, False, torch.float16, torch.float16, torch.float16 + torch.float16, False, torch.float16, torch.float16, torch.float16, "matmul" ) def test_flag_false_bf16_path(self): self._run_case( - torch.bfloat16, False, torch.bfloat16, torch.bfloat16, torch.bfloat16 + torch.bfloat16, + False, + torch.bfloat16, + torch.bfloat16, + torch.bfloat16, + "matmul", )