flashinfer.padded_seq_len

flashinfer.padded_seq_len(max_seq_len: int) → int

Return the minimum column count of the paged MQA logits out buffer.

Rounds max_seq_len up to the kernels’ internal output-store granularity. The kernels write the padded trailing positions unconditionally, so the output tensor must be allocated with at least this many columns (a narrower out is an out-of-bounds write). The granularity is an implementation detail that may change – always call this helper rather than hard-coding the padding. The padding columns hold unspecified scratch – never consume them.

Use this to pre-allocate the out parameter:
out = torch.empty(

(B * next_n, padded_seq_len(max_seq_len)), dtype=…, device=”cuda”

) logits = fp8_paged_mqa_logits(…, out=out)