How FlashInfer KV Cache Attention Accelerates Streaming 3D Reconstruction Over SDPA
FlashInfer KV cache attention accelerates streaming 3D reconstruction by replacing dense tensor concatenation with a paged memory layout, reducing per-frame complexity from O(F) to O(1) while eliminating repeated kernel launches through single-plan reuse.
Streaming 3D reconstruction requires each new camera frame to attend to a growing history of past frames while keeping memory and compute bounded. In the Robbyant/lingbot-map repository, the FlashInferKVCacheManager and FlashInferAttention classes implement a paged attention mechanism that outperforms the baseline scaled dot-product attention (SDPA) by decoupling storage capacity from sequence length.
Performance Bottlenecks in SDPA-Based Streaming
Dense Memory Layout and Linear Scaling
The baseline CausalAttention and SDPAAttention classes in lingbot_map/layers/attention.py store all past key/value (K/V) tensors in a Python dictionary and concatenate them into a dense tensor for every forward pass. This approach consumes O(N·F) memory where F is the number of frames, forcing the model to copy increasingly large tensors as the streaming session progresses.
Repeated Kernel Launch Overhead
Because SDPA operates on the full concatenated sequence, every incoming frame triggers a new kernel launch with a different sequence length. This incurs repeated CUDA kernel compilation and launch overhead, significantly increasing latency in the streaming loop compared to optimized paged alternatives.
Manual Cache Eviction Complexity
Evicting old frames in the SDPA baseline requires manual tensor slicing, data copying, and mask reconstruction. These operations block the CUDA stream and consume additional memory bandwidth, creating performance bottlenecks during real-time reconstruction.
Architectural Advantages of FlashInfer KV Cache Attention
Paged KV Cache Memory Layout
The FlashInferKVCacheManager in lingbot_map/layers/flashinfer_cache.py pre-allocates a fixed tensor of shape [max_num_pages, 2, page_size, H, D]. Instead of concatenating tensors, the system writes frames into fixed-size pages and tracks them via indices. This bounds memory usage regardless of sequence length, reducing the per-frame memory cost to O(1).
Two-Stream Design for Token Types
The implementation maintains two distinct page streams:
- Patch pages: Recyclable storage for the sliding window of recent frames, managed via deques (
free_patch_pages,scale_patch_pages,live_window_patch_pages). - Special pages: Append-only storage for camera tokens, register tokens, and scale tokens that must never be evicted.
This separation allows eviction to occur by simply moving page IDs between deques without copying tensor data.
Plan Reuse and Kernel Optimization
The FlashInferAttention class orchestrates the forward pass using BatchPrefillWithPagedKVCacheWrapper. It builds the attention plan once per frame step when block_idx == 0, then reuses this plan for all subsequent transformer layers. This eliminates per-layer kernel launch costs, contrasting sharply with SDPA which pays these costs for every layer and every frame.
Efficient Attention Masking and Padding
Special tokens are placed last in the page table, allowing the kernel to use paged_kv_last_page_len to describe the partial tail page. This eliminates the need for custom boolean masks. When fa3=False, the page size is set exactly to patches_per_frame to eliminate padding waste; when fa3=True requires power-of-two alignment, padding is isolated to the special stream only.
Implementation in the lingbot-map Codebase
FlashInferKVCacheManager Responsibilities
The manager exposes a drop-in API with four key methods:
append_frame: Writes current frame K/V tensors into the paged cache.evict_frames: Recycles oldest patch pages while preserving required scale and sliding-window pages.compute_attention: Executes the FlashInfer kernel using the pre-built plan.reset: Clears the cache state between sequences.
FlashInferAttention Forward Pass
In lingbot_map/layers/attention.py, the FlashInferAttention.forward method executes the streaming logic:
- Converts Q/K/V to NHD layout via
prepare_qkv. - Calls
manager.append_frameto write the current frame. - Calls
manager.evict_framesto maintain the sliding window (kv_cache_sliding_window) and scale frame (kv_cache_scale_frames) constraints. - Calls
manager.compute_attentionwith the shared plan to execute attention.
In batch mode (no KV cache), the layer falls back to regular SDPA for numerical consistency.
Integration with Streaming Models
The GCTStream class and related window variants in lingbot_map/models/gct_stream.py integrate these layers into the full reconstruction pipeline, passing the kv_cache parameter as the manager instance itself.
Practical Code Examples
Streaming Reconstruction Loop with FlashInfer
# ----------------------------------------------------------------------
# Example: Using FlashInferAttention in a streaming reconstruction loop
# ----------------------------------------------------------------------
import torch
from lingbot_map.layers.attention import FlashInferAttention
from lingbot_map.layers.flashinfer_cache import FlashInferKVCacheManager
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
batch = 1 # streaming uses batch‑size = 1
tokens_per_frame = 262 # 256 patches + 6 special tokens
num_blocks = 12 # number of transformer blocks / layers
num_heads = 16
head_dim = 64
# Initialise the paged KV cache manager (once for the whole model)
kv_manager = FlashInferKVCacheManager(
num_blocks=num_blocks,
max_num_frames=88, # scale + window + headroom
tokens_per_frame=tokens_per_frame,
num_heads=num_heads,
head_dim=head_dim,
dtype=torch.bfloat16,
device=device,
)
# Initialise a FlashInfer‑enabled attention layer
attn = FlashInferAttention(
dim=num_heads * head_dim,
num_heads=num_heads,
rope=None, # optional RoPE module
kv_cache_sliding_window=64,
kv_cache_scale_frames=8,
).to(device)
# Simulated streaming loop
for frame_idx in range(200): # 200 incoming frames
# x: [B, N, C] (N = tokens_per_frame, C = dim)
x = torch.randn(batch, tokens_per_frame, num_heads * head_dim,
dtype=torch.bfloat16, device=device)
# Forward pass with the manager – note `kv_cache` is the manager itself
out = attn(
x,
num_frames=1, # single‑frame streaming
kv_cache=kv_manager,
global_idx=0, # block index for this layer
)
# `out` now contains the attention‑processed tokens for the current frame
Comparing SDPA vs. FlashInfer
# ----------------------------------------------------------------------
# Example: Comparing SDPA vs. FlashInfer for a single block
# ----------------------------------------------------------------------
import torch
from lingbot_map.layers.attention import SDPAAttention, FlashInferAttention
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
tokens = 262
B, C = 1, 192 # dim = 8 heads × 24 d‑model (example)
x = torch.randn(B, tokens, C, device=device, dtype=torch.bfloat16)
# SDPA baseline (no KV cache)
sdpa = SDPAAttention(dim=C, num_heads=8).to(device)
out_sdpa = sdpa(x, kv_cache=None) # full dense attention each frame
# FlashInfer streaming (with paged KV cache)
kv_mgr = FlashInferKVCacheManager(
num_blocks=1,
max_num_frames=88,
tokens_per_frame=tokens,
num_heads=8,
head_dim=C // 8,
dtype=torch.bfloat16,
device=device,
)
flash = FlashInferAttention(dim=C, num_heads=8).to(device)
out_flash = flash(x, kv_cache=kv_mgr, num_frames=1, global_idx=0)
print("SDPA shape:", out_sdpa.shape, "FlashInfer shape:", out_flash.shape)
Summary
- FlashInfer KV cache attention replaces dense tensor concatenation with a paged memory layout, reducing per-frame memory complexity from O(F) to O(1).
- The two-stream design separates recyclable patch pages from append-only special pages, enabling zero-copy eviction via deque manipulation.
- Plan reuse eliminates per-layer kernel launches by building the
BatchPrefillWithPagedKVCacheWrapperplan once per frame step and sharing it across all transformer blocks. - Optimized memory alignment minimizes padding waste by matching page size to
patches_per_framewhen possible, or isolating padding to the special token stream. - Implementation resides in
lingbot_map/layers/flashinfer_cache.pyandlingbot_map/layers/attention.py, providing a drop-in replacement forSDPAAttention.
Frequently Asked Questions
What is the main performance advantage of FlashInfer KV cache attention over SDPA?
FlashInfer KV cache attention reduces the computational complexity per frame from O(F) to O(1) by using a pre-allocated paged cache instead of concatenating all past frames. It also eliminates repeated CUDA kernel launches through plan reuse, significantly reducing latency in streaming 3D reconstruction pipelines.
How does the two-stream page design handle cache eviction?
The FlashInferKVCacheManager maintains separate deques for patch pages (recyclable) and special pages (append-only). Eviction occurs by moving page IDs between live_window_patch_pages and free_patch_pages without copying tensor data, allowing constant-time memory management regardless of frame history length.
Why is plan reuse important for transformer inference?
Building a BatchPrefillWithPagedKVCacheWrapper plan involves calculating page table indices and kernel parameters. By constructing this plan once when block_idx == 0 and reusing it for all subsequent layers in the same frame, FlashInfer avoids the overhead of repeated kernel compilation and launch that plagues per-layer SDPA implementations.
Can FlashInferAttention fall back to SDPA for non-streaming use cases?
Yes. The FlashInferAttention class in lingbot_map/layers/attention.py automatically detects when kv_cache=None and falls back to standard SDPA for batch processing. This ensures numerical consistency between streaming and batch modes while allowing the same model weights to serve both use cases.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →