flashinfer.compute_paged_mqa_logits_schedule

flashinfer.compute_paged_mqa_logits_schedule(seq_lens: Tensor, device: device | None = None, *, next_n: int = 1, variant: str = 'fp8', use_gpu_kernel: bool = True, out: Tensor | None = None) → Tensor

Compute the CTA schedule tensor for paged MQA logits kernels.

Returns [num_sms+1, 2] int32 on CUDA, ready to pass as schedule_meta to fp8_paged_mqa_logits() or fp4_paged_mqa_logits(). Pass the same variant and next_n as the upcoming main call: for some (variant, next_n, device) combinations the kernel internally restructures the problem, and this helper applies the same internal policy so the returned schedule always describes the work the kernel actually runs. Treat the contents as opaque – never inspect, serialize, or reuse them across flashinfer versions, devices, or different (variant, next_n).

Parameters:
  • seq_lens – [B] int32 (1-D), on CPU or on device; any stride.

  • device – target CUDA device. Defaults to seq_lens.device, or the current CUDA device when seq_lens is on CPU.

  • next_n – the next_n (q.shape[1]) of the upcoming main call. Defaults to 1 (plain decode).

  • variant – “fp8” or “fp4” – which paged-MQA API the schedule is for. Defaults to “fp8”.

  • use_gpu_kernel – if True (default), compute entirely on-GPU via a small dedicated schedule kernel – no D2H copy, CUDA-graph-capturable. Falls back to CPU numpy (not graph-capturable) when the CuTe DSL is unavailable, cannot target the device’s architecture, or seq_lens is on CPU.

  • out – optional pre-allocated [num_sms+1, 2] int32 on CUDA. Required for CUDA-graph capture (static buffer). If None, a new tensor is allocated each call. This provides static storage, not static contents: the address stays stable, but the values must be recomputed into it whenever the contents of seq_lens change – including on every graph replay where the lengths may have moved.

Returns:

[num_sms+1, 2] int32 on CUDA (out if provided).

Return type:

schedule_meta