sgl-project/sglang · error · TypeError
LSE tensor must be Float32
Error message
LSE tensor must be Float32
What it means
The optional final LSE tensor passed to the FA4 combine kernel must be float32 (or omitted). Like the partials, it holds log-sum-exp accumulators that require fp32 precision.
Source
Thrown at python/sglang/kernels/ops/attention/flash_attn/cute/flash_fwd_combine.py:227
mO: cute.Tensor,
mLSE: Optional[cute.Tensor] = None,
cu_seqlens: Optional[cute.Tensor] = None,
seqused: Optional[cute.Tensor] = None,
num_splits_dynamic_ptr: Optional[cute.Tensor] = None,
varlen_batch_idx: Optional[cute.Tensor] = None,
semaphore_to_reset: Optional[cute.Tensor] = None,
# Always keep stream as the last parameter (EnvStream: obtained implicitly via TVM FFI).
stream: cuda.CUstream = None,
):
# Type checking
if const_expr(not (mO_partial.element_type == self.dtype_partial)):
raise TypeError("O partial tensor must match dtype_partial")
if const_expr(not (mO.element_type == self.dtype)):
raise TypeError("O tensor must match dtype")
if const_expr(mLSE_partial.element_type not in [Float32]):
raise TypeError("LSE partial tensor must be Float32")
if const_expr(mLSE is not None and mLSE.element_type not in [Float32]):
raise TypeError("LSE tensor must be Float32")
# Shape validation - input tensors are in user format, need to be converted to kernel format
if const_expr(len(mO_partial.shape) not in [4, 5]):
raise ValueError(
"O partial tensor must have 4 or 5 dimensions: (num_splits, batch, seqlen, nheads, headdim) or (num_splits, total_q, nheads, headdim)"
)
if const_expr(len(mLSE_partial.shape) not in [3, 4]):
raise ValueError(
"LSE partial tensor must have 3 or 4 dimensions: (num_splits, batch, seqlen, nheads) or (num_splits, total_q, nheads)"
)
if const_expr(len(mO.shape) not in [3, 4]):
raise ValueError(
"O tensor must have 3 or 4 dimensions: (batch, seqlen, nheads, headdim) or (total_q, nheads, headdim)"
)
if const_expr(mLSE is not None and len(mLSE.shape) not in [2, 3]):
raise ValueError(
"LSE tensor must have 2 or 3 dimensions: (batch, seqlen, nheads) or (total_q, nheads)"
)View on GitHub (pinned to 0132848349)
Solutions
- Allocate mLSE as float32 or pass None if you don't need LSE output
- Check LSE dtype at the call site before invoking the combine
Example fix
// before mLSE = torch.empty(lse_shape, dtype=torch.bfloat16) // after mLSE = torch.empty(lse_shape, dtype=torch.float32)
Defensive patterns
Strategy: validation
Validate before calling
if mLSE is not None:
assert mLSE.dtype == torch.float32 Type guard
def lse_ok(t) -> bool: return t is None or t.dtype == torch.float32
Prevention
- Pass mLSE=None when LSE output isn't consumed
- Never reuse model-dtype buffers for LSE
When it happens
Trigger: Passing mLSE (non-None) with element type other than Float32 to the combine __call__.
Common situations: Reusing a bf16 buffer for LSE output, or passing an LSE tensor produced by a different library with a different dtype convention.
Related errors
- O tensor must match dtype
- LSE partial tensor must be Float32
- v_cache must be provided
- q can only be None when only_qv=True
- q must be provided unless qv is provided with only_qv=True
AI-assisted analysis of sgl-project/sglang@0132848349 (2026-08-28).
Data as JSON: /api/errors/387d509bd9db1a60.
Report an issue: GitHub.