vllm.v1.attention.backends.mla.flashinfer_mla_sparse_sm90
¶
FlashInfer sparse MLA backend for SM90 (Hopper) NoPE models.
Wraps FlashInfer's BatchMLAPagedAttentionWrapper (FA2/FA3 paths), which
as of FlashInfer 0.6.18 supports head_dim_kpe=0 (GLM-5.3-Flash NoPE MLA)
and FP8 E4M3 KV caches on SM90 with in-kernel dequantization: the FP8 cache
is read directly (half the bf16 HBM traffic) and converted to BF16 in shared
memory, while queries stay BF16 (no query quantization).
Sparsity rides the same trick the FA-based sparse backend uses: with
page_size=1 the per-token top-k slot indices ARE the page table, so each
query token becomes one varlen batch row whose kv_indices slice is its
top-k row and whose kv_len is its valid count. Causality is already
encoded by the indexer's selection, so causal=False.
CUDA-graph handling: plan() copies its inputs to host unconditionally,
so it must stay outside graph capture. Each metadata builder owns a wrapper,
reserved capture-stable device buffers, and the plan parameters. The wrapper
bakes the per-row kv_len into its int schedule at plan() time — run()
never reads the device-side buffer. The builder plans (outside capture) from
sync-free host upper bounds and clamps each scheduled work item's kv_end
on device to the exact valid count, so the kernel never reads past a row's
valid prefix (the -1 tail of the converted index buffer) and async scheduling
needs no D2H sync. Per-step content (top-k slots) is written into the
reserved buffers by kernels inside the captured forward, and captured runs
read the refreshed plan buffers on replay.
KV cache format: plain contiguous E4M3 [num_blocks, block_size, 512]
(uint8 storage) with a per-tensor k_scale; BF16 caches also work. The
per-token x 128-channel-group ckv_scale_arr layout is supported by the
kernel but not wired yet (it needs a group-quantizing cache-write op).
Classes:
-
FlashInferMLASparseSM90Builder–Reuse the common sparse metadata (req ids, topk buffer access).
FlashInferMLASparseSM90Builder
¶
Bases: FlashInferMLASparseMetadataBuilder
Reuse the common sparse metadata (req ids, topk buffer access).
Source code in vllm/v1/attention/backends/mla/flashinfer_mla_sparse_sm90.py
419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 | |
_kv_lens_host(cam)
¶
Host per-row KV lengths for the plan and whether they are exact.
A row for the j-th query token of request i attends
seq_lens[i] - q_len[i] + j + 1 tokens. The indexer's selection
then bounds the valid count: contexts up to index_topk select
everything (valid == context); longer contexts keep the top
index_topk pool-expanded tokens plus the trailing incomplete
pool (valid == index_topk + context % index_kpool). Both match
the count of non -1 entries the convert kernel produces.
With seq_lens_cpu_upper_bound (optimistic under async spec
decode) returns sync-free upper bounds, which the plan clamps on
device. Otherwise returns exact lengths via a D2H copy.
Source code in vllm/v1/attention/backends/mla/flashinfer_mla_sparse_sm90.py
_SM90State
¶
Builder-owned wrapper, capture-stable buffers, and plan parameters.
One instance serves every MLA layer in an attention group because the plan depends only on the batch shape, not the layer.
Methods:
-
plan–Plan per-row KV lengths (CPU int32,
[num_tokens]).
Source code in vllm/v1/attention/backends/mla/flashinfer_mla_sparse_sm90.py
230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 | |
plan(num_tokens, kv_lens, cam, req_id_per_token)
¶
Plan per-row KV lengths (CPU int32, [num_tokens]).
The wrapper bakes kv_len into its int schedule from host values, so
rows past their valid count would send the kernel into the -1 tail of
the converted index buffer. With cam and req_id_per_token
(device batch layout), kv_lens only needs to upper-bound the
valid counts: the plan is padded by _PLAN_SLACK and reused while
the bound still fits, and every call clamps each work item's kv_end
on device to the exact count, so no D2H sync is needed. Without them
the lengths must be exact. Must run outside CUDA graph capture: the
in-place refreshed plan_info/indptr buffers are what captured runs
read.