How to Improve Streaming Inference Performance with FlashInfer KV Cache Backend
The Ling-Bot MAP library achieves high-throughput streaming inference by using a two-stream paged attention design that separates recyclable patch pages from append-only special tokens, enabling single-allocation memory management and plan-once-run-many execution across transformer blocks.
The Ling-Bot MAP repository implements a specialized KV-cache manager leveraging FlashInfer kernels for low-latency streaming scenarios where frames arrive sequentially. By optimizing the paging strategy and attention planning logic in lingbot_map/layers/flashinfer_cache.py, the library minimizes host-side overhead and maximizes GPU utilization for real-time video and audio processing applications.
Two-Stream Page Pool Architecture
At the core of the performance optimization lies a two-stream page pool design implemented in lingbot_map/layers/flashinfer_cache.py (lines 19-22). This architecture maintains separate memory pools for different token types:
- Patch pages: Recyclable storage for sliding window frames that fall outside the attention scale
- Special pages: Append-only storage for 6 special tokens per frame (camera, registers, scale)
This separation guarantees that expensive CUDA allocation occurs only once during construction. Subsequent frames merely move page IDs between deques rather than triggering new memory allocations. The physical layout pre-allocates a tensor of shape [max_num_pages, 2, page_size, H, D] (lines 41-48), ensuring all pages are ready for immediate use.
Critical Performance Components
Understanding the specific components that govern speed allows precise tuning of the FlashInfer backend.
Page-Size Handling for FA2 vs FA3
The manager adapts page sizes based on the FlashAttention version. In lingbot_map/layers/flashinfer_cache.py (lines 99-105), the logic distinguishes between:
- FA2 (default): Supports exact page sizes matching the number of patches per frame, eliminating zero-padding and reducing memory traffic
- FA3 (SM90): Requires power-of-two page sizes, necessary only for H100-class GPUs exploiting newer kernels
For most deployments, keeping fa3=False maximizes throughput by avoiding wasted padding bytes.
Lazy KV-Cache Manager Initialization
To prevent unnecessary overhead when running in standard SDPA mode, the manager employs lazy initialization. In lingbot_map/aggregator/stream.py (lines 198-215), the _get_flashinfer_manager method constructs the cache only upon first frame processing. This avoids allocating massive tensors when the model operates in non-streaming contexts.
Deferred Eviction Strategy
The FlashInferKVCacheManager supports deferred eviction through the _defer_eviction flag (lines 72-78). When enabled, the system can delay recycling patch pages via execute_deferred_eviction or rollback frames using rollback_last_frame (lines 56-64). This capability proves essential when keyframe decisions remain pending, preventing premature page recycling during uncertain states.
One-Time Planning Pattern
The attention implementation follows a plan-once-run-many paradigm. In compute_attention (lines 86-115), the plan() method executes only on block 0 to build the visible page table (scale → window → special) and calculate the last page length. Subsequent transformer blocks reuse this plan via run(), eliminating repeated host-to-device planning work across all layers.
Persistent Workspace Pre-allocation
The FlashInfer wrapper pre-allocates a 128 MiB persistent workspace buffer during initialization (lines 184-191). This buffer serves all attention calls, avoiding the latency spikes associated with repeated CUDA memory allocations during inference.
Force-FP32 Fallback Avoidance
When input dtype is float32, the manager cannot use FlashInfer's FA2 kernels. Instead, it gathers K/V into dense tensors and executes scaled_dot_product_attention (lines 72-84). This fallback path significantly reduces throughput, making it critical to maintain bfloat16 or float16 precision throughout the model.
Practical Performance Tuning
Implement these specific optimizations to extract maximum throughput from the FlashInfer backend:
- Maintain low-precision dtypes: Keep KV tensors in
bfloat16orfloat16to avoid the slower FP32 gather-and-SDPA fallback path - Use exact page sizes: Set
fa3=Falseto eliminate zero-padding in patch pages unless running on H100 hardware - Size frames conservatively: Increase
max_total_framesonly as needed, as larger values expand the special-page pool (line 33) and consume GPU memory - Batch planning when possible: Accumulate multiple frames before calling
plan()to reduce host-side planning frequency, provided latency constraints permit - Profile synchronization points: Use
torch.cuda.profilerandtorch.cuda.synchronize()aroundcompute_attentionto identify hidden bottlenecks
Code Examples
Minimal Streaming Inference Setup
Configure the FlashInferKVCacheManager directly for fine-grained control over the streaming process:
import torch
from lingbot_map.layers.flashinfer_cache import FlashInferKVCacheManager
device = torch.device("cuda")
tokens_per_frame = 262 # 256 patches + 6 special tokens
num_blocks = 12 # number of transformer blocks
num_heads = 16
head_dim = 64
dtype = torch.bfloat16 # keep in bfloat16 for Flash-Infer FA2
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=dtype,
device=device,
num_special_tokens=6,
scale_frames=8,
sliding_window=64,
max_total_frames=200,
force_fp32=False, # stay in fp16/bf16
fa3=False, # use FA2 (exact page size)
)
# Simulate a stream of frames
for frame_idx in range(100):
# K/V tensors for this frame (shape [tokens_per_frame, H, D])
k = torch.randn(tokens_per_frame, num_heads, head_dim, dtype=dtype, device=device)
v = torch.randn(tokens_per_frame, num_heads, head_dim, dtype=dtype, device=device)
# Append to *all* blocks (for simplicity we just use block 0 here)
kv_manager.append_frame(0, k, v)
# Evict old patch pages (scale & window are preserved automatically)
kv_manager.evict_frames(0, scale_frames=8, sliding_window=64)
# When you need attention on the current frame:
q = torch.randn(tokens_per_frame, num_heads, head_dim, dtype=dtype, device=device)
out = kv_manager.compute_attention(0, q) # triggers plan() once, then run()
High-Level AggregatorStream Interface
For production deployments, use the AggregatorStream class which handles manager lifecycle automatically:
from lingbot_map.aggregator.stream import AggregatorStream
# Build the model (only the streaming part is shown)
stream = AggregatorStream(
depth=12,
num_heads=16,
head_dim=64,
tokens_per_frame=262,
use_sdpa=False, # Flash-Infer is the default backend
device=torch.device("cuda"),
dtype=torch.bfloat16,
)
# Forward a single frame (features already tokenised)
features = torch.randn(1, 262, 16, 64, device="cuda", dtype=torch.bfloat16)
output = stream(features, causal_inference=False) # runs Flash-Infer path
Deferred Eviction for Keyframe Decisions
Enable deferred eviction when frame validity remains uncertain:
# Enable deferred eviction
kv_manager._defer_eviction = True
# Append a few frames without evicting
for _ in range(5):
k = torch.randn(tokens_per_frame, num_heads, head_dim, device=device, dtype=dtype)
v = torch.randn(tokens_per_frame, num_heads, head_dim, device=device, dtype=dtype)
kv_manager.append_frame(0, k, v)
# When the decision is ready, execute eviction
kv_manager.execute_deferred_eviction(0, scale_frames=8, sliding_window=64)
# Or rollback if the last frame should be discarded
kv_manager.rollback_last_frame(0)
Summary
- The two-stream page pool architecture in
lingbot_map/layers/flashinfer_cache.pyseparates recyclable patch pages from append-only special tokens, eliminating runtime allocations - Plan-once-run-many execution across transformer blocks minimizes host-side overhead by computing the visible page table only on block 0
- Lazy initialization in
lingbot_map/aggregator/stream.pyprevents unnecessary memory allocation when operating in non-streaming modes - Deferred eviction capabilities support pending keyframe decisions without forcing premature cache recycling
- Maintaining bfloat16/float16 precision avoids the slower Force-FP32 fallback path that bypasses optimized FlashInfer kernels
Frequently Asked Questions
What is the optimal page size configuration for FlashInfer in Ling-Bot MAP?
For FlashAttention 2 (FA2), which is the default backend, configure the page size to match the exact number of patches per frame by setting fa3=False. This eliminates zero-padding and reduces memory traffic compared to the power-of-two requirements of FA3. Only enable FA3 when running on H100-class GPUs (SM90) that can exploit the newer kernels, as implemented in lingbot_map/layers/flashinfer_cache.py (lines 99-105).
How does the deferred eviction mechanism improve streaming performance?
The deferred eviction mechanism allows the KV-cache manager to delay recycling patch pages until a keyframe decision is finalized. By setting _defer_eviction = True, the system accumulates frames without evicting, then either executes execute_deferred_eviction to commit or rollback_last_frame to discard. This prevents unnecessary page recycles when frame validity remains uncertain, optimizing cache utilization in lingbot_map/layers/flashinfer_cache.py (lines 56-78).
Why should I avoid float32 precision when using the FlashInfer backend?
When KV tensors use float32 dtype, FlashInfer's optimized FA2 kernels cannot execute. The manager falls back to gathering K/V into dense tensors and running PyTorch's scaled_dot_product_attention, which significantly reduces throughput. The fallback path is implemented in compute_attention (lines 72-84), making it critical to maintain bfloat16 or float16 precision throughout the model for maximum performance.
How does lazy initialization reduce overhead in non-streaming scenarios?
The AggregatorStream class creates the FlashInferKVCacheManager only when the first frame arrives via _get_flashinfer_manager (lines 198-215 in lingbot_map/aggregator/stream.py). This prevents allocating the large persistent workspace buffer and page pools when the model runs in SDPA mode or processes non-streaming inputs, ensuring resources are consumed only when FlashInfer capabilities are actually required.
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 →