sgl-project/sglang · error · TypeError
O tensor must match dtype
Error message
O tensor must match dtype
What it means
The FlashAttention-4 CuTe split-K combine kernel requires the output tensor mO to have the exact element dtype the combine object was constructed with (self.dtype). A dtype mismatch means the kernel would write garbage or fail to compile, so it is rejected up front.
Source
Thrown at python/sglang/kernels/ops/attention/flash_attn/cute/flash_fwd_combine.py:223
def __call__(
self,
mO_partial: cute.Tensor,
mLSE_partial: cute.Tensor,
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)"
)View on GitHub (pinned to 0132848349)
Solutions
- Allocate/convert mO to match the combine object's dtype: mO = torch.empty(..., dtype=combine.dtype, device=...)
- Check where the combine object is constructed and align its dtype argument with your output tensor
- If bridging precisions, do an explicit .to(dtype) after the kernel instead of passing a mismatched buffer
Example fix
// before mO = torch.empty(shape, dtype=torch.float16) combine(mO_partial, mLSE_partial, mO, mLSE) // after mO = torch.empty(shape, dtype=combine.dtype) combine(mO_partial, mLSE_partial, mO, mLSE)
Defensive patterns
Strategy: type-guard
Validate before calling
assert mO.dtype == torch_dtype_used_for_combine, f'expected {combine.dtype}, got {mO.dtype}' Type guard
def out_matches_combine(mO, combine) -> bool:
return mO.dtype == combine.dtype Prevention
- Allocate the output tensor from the combine object's dtype, not a hardcoded torch dtype
- Centralize dtype selection in one config constant used by both quantizer and buffers
When it happens
Trigger: Calling FlashAttnCombineCombine (or the combine stage of a split-K FA4 run) with an mO tensor whose dtype differs from the dtype the combine class was instantiated with, e.g. bf16 combine object but fp16 output tensor.
Common situations: Mixed-precision attention setups: partials accumulated in one dtype (e.g. fp16) but the final output buffer allocated as bf16, or reusing a combine object across models with different dtypes.
Related errors
- LSE partial tensor must be Float32
- LSE 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/61439f2963f5095a.
Report an issue: GitHub.