sgl-project/sglang · error · ValueError

LSE tensor must have 2 or 3 dimensions: (batch, seqlen, nhea

Error message

LSE tensor must have 2 or 3 dimensions: (batch, seqlen, nheads) or (total_q, nheads)

What it means

The optional final LSE tensor must be 2D (batch, nheads) batched or 3D (total_q, nheads) varlen — again one rank below the LSE partials.

Source

Thrown at python/sglang/kernels/ops/attention/flash_attn/cute/flash_fwd_combine.py:243

            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)"
            )

        mO_partial, mO = [assume_tensor_aligned(t) for t in (mO_partial, mO)]
        # (num_splits, b, seqlen, h, d) -> (seqlen, d, num_splits, h, b)
        # or (num_splits, total_q, h, d) -> (total_q, d, num_splits, h)
        O_partial_layout_transpose = (
            [2, 4, 0, 3, 1] if const_expr(cu_seqlens is None) else [1, 3, 0, 2]
        )
        # (b, seqlen, h, d) -> (seqlen, d, h, b) or (total_q, h, d) -> (total_q, d, h)
        mO_partial = cute.make_tensor(
            mO_partial.iterator,
            cute.select(mO_partial.layout, mode=O_partial_layout_transpose),
        )
        O_layout_transpose = (
            [1, 3, 2, 0] if const_expr(cu_seqlens is None) else [0, 2, 1]
        )
        mO = cute.make_tensor(

View on GitHub (pinned to 0132848349)

Solutions

  1. Batched: mLSE shape (batch, nheads); varlen: (total_q, nheads)
  2. Or pass mLSE=None if LSE output is not needed

Example fix

// before
mLSE = torch.empty(num_splits, b, h)
// after
mLSE = torch.empty(b, h, dtype=torch.float32)
Defensive patterns

Strategy: validation

Validate before calling

if mLSE is not None:
    assert mLSE.dim() in (2, 3) and mLSE.dim() == mLSE_partial.dim() - 1

Type guard

def final_lse_ok(lse, lse_partial) -> bool:
    return lse is None or (lse.dim() in (2, 3) and lse.dim() == lse_partial.dim() - 1)

Prevention

When it happens

Trigger: Passing mLSE (non-None) with rank other than 2 or 3.

Common situations: Carrying over the num_splits axis, or allocating LSE with an extra seqlen dim in varlen mode.

Related errors


AI-assisted analysis of sgl-project/sglang@0132848349 (2026-08-28). Data as JSON: /api/errors/af5c5539afdad2f3. Report an issue: GitHub.