flashinfer.precompile_paged_mqa_logits

flashinfer.precompile_paged_mqa_logits(device: device | None = None, variants: Tuple[str, ...] = ('fp8', 'fp4'), output_dtypes: Tuple[dtype, ...] | None = None, batch_sizes: Sequence[int] | None = None) → None

Pre-compile paged MQA logits kernels for common static configs.

Populates the on-disk CuTe-DSL kernel cache so subsequent calls to fp8_paged_mqa_logits() and fp4_paged_mqa_logits() skip compilation on first use. Call once during deployment setup or as part of a package-build step.

The warmed set is exactly: fp8 – num_heads=64, head_dim=128, block_size in {32, 64, 128}, next_n 1..4, fp32 epilogue/accumulator (12 kernels); fp4 – num_heads=64, head_dim=128, block_size in {32, 64, 128}, next_n 1..4, fp32 epilogue (12 kernels per output dtype). Anything else (fp16 epilogue, other head geometry) still compiles on first use. Measured on sm_100a: ~1s per fp8 kernel, ~3s per output dtype for the 12 fp4. The fp4 next_n=4 entry compiles the decomposition the fixed policy picks on the target device (direct on Rubin, two internal passes on Blackwell).

Parameters:
  • device – CUDA device to target. Defaults to the current CUDA device, so a worker that has set its per-rank device gets kernels for that device without passing anything.

  • variants – Which precisions to build. A deployment normally runs one indexer precision, so pass e.g. ("fp8",) to avoid compiling kernels that will never be called.

  • output_dtypes – Which output dtypes to warm. output_dtype is part of the compile cache key, so a dtype not warmed here still compiles on first use. Defaults to each variant’s common set: float32 for FP8, and both bfloat16 and float32 for FP4 – the API default plus the dtype consumers with a float logits ABI require. Pass an explicit tuple to build only what you run.

  • batch_sizes – Batch sizes whose GPU schedule kernel should be warmed, in caller units. The schedule kernel specialises on the scheduler row count in 32-row buckets: the batch size for fp8 and unsplit fp4, but batch_size*num_atoms under the fp4 atom split (next_n=4 on SM100/SM103 schedules twice the rows). Every bucket the requested variants can reach from each size is warmed, then deduplicated; none of this is covered by the shape sweep above. Defaults to None, which warms no schedule buckets – a deployment that captures CUDA graphs for a known set of batch sizes should pass them, or the first capture of each bucket pays compilation.