How to Implement Long Context Inference (Up to 139K Tokens) with KTransformers
You can run inference on sequences up to 139,000 tokens in KTransformers by enabling mode="long_context", which splits the KV-cache into dynamic blocks and offloads inactive blocks to CPU or disk while maintaining a local GPU window.
KTransformers (kvcache-ai/ktransformers) extends transformer inference beyond native context limits through a block-wise KV-cache architecture. By dynamically selecting which cache blocks reside in GPU memory and spilling the rest to tiered storage, the framework enables single-GPU inference on contexts approaching 139K tokens without exhausting VRAM.
Understanding the Block-Wise KV-Cache Architecture
Long-context inference in KTransformers relies on three core components working in concert:
Config.long_context_config– Parsed from your YAML configuration inarchive/ktransformers/server/config/config.py(lines 174-188), this defines maximum sequence length, block granularity, and memory policies.DynamicScaledDotProductAttention– Located inarchive/ktransformers/operators/dynamic_attention.py(lines 30-53), this attention engine manages block allocation, pre-selection tables, and kernel selection (Flash Attention or CPU fallback).- Utility wrappers – Functions
chunk_prefillanddecode_wrapperinarchive/ktransformers/util/utils.py(lines 496-528) handle the orchestration of chunked prefill and token generation.
The system stores KV tensors in cache_key_states and cache_value_states data structures (initialized at lines 101-115 of dynamic_attention.py) that can reside on GPU, CPU, or disk depending on access patterns.
Step 1: Configure the YAML for Extended Context
Create a configuration file that specifies your target sequence length and block management strategy. Set max_seq_len up to 139000 tokens according to your hardware constraints.
# config.yaml
model:
name_or_path: "meta-llama/Llama-2-70b"
device: "cuda:0"
long_context:
max_seq_len: 139000 # Target context length (139K tokens)
block_size: 128 # Tokens per KV-cache block
local_windows_len: 4096 # Sliding window kept resident on GPU
second_select_num: 32 # Top-k blocks retained per layer
anchor_type: "DYNAMIC" # Block selection policy
kv_type: "FP16" # Cache precision (FP16, FP32, Q4_0, Q8_0)
dense_layer_num: 2 # Layers exempt from offloading
preselect_block: true # Enable fast pre-selection table
head_select_mode: "SHARED" # Shared block mask across heads
preselect_block_count: 96 # Blocks maintained in pre-selector
layer_step: 1 # Stride for layer scanning
token_step: 100 # Stride for token scanning within blocks
chunk_size: 20480 # Prefill chunk size
The block_size parameter determines memory granularity—128 tokens per block is optimal for Llama-style architectures. The local_windows_len defines how many recent tokens remain in high-speed GPU memory, while older blocks migrate to slower tiers.
Step 2: Load the Model in Long-Context Mode
Instantiate the configuration and load your model with the mode parameter set to "long_context". This triggers the block-wise attention pathway in the utility loader.
from ktransformers.server.config.config import Config
from ktransformers.util.utils import load_model
# Load global configuration
cfg = Config() # Automatically reads ~/.ktransformers/config.yaml
assert cfg.long_context_config, "Long context configuration required"
# Initialize model with block-wise attention engine
model, tokenizer = load_model(
cfg.model["name_or_path"],
device=cfg.model["device"],
mode="long_context", # Critical switch for block-wise caching
)
When mode="long_context" is active, utils.py routes embeddings through CPU staging (inputs_embeds = model.model.embed_tokens(inputs.to("cpu"))) and initializes the tiered KV-cache structure.
Step 3: Initialize the Dynamic Attention Engine
For custom inference loops, instantiate DynamicScaledDotProductAttention directly using parameters from your configuration. This class manages the cache_importance tensor and block selection logic.
from transformers import AutoConfig
import torch
from ktransformers.operators.dynamic_attention import DynamicScaledDotProductAttention
cfg_hf = AutoConfig.from_pretrained(cfg.model["name_or_path"])
device = torch.device(cfg.model["device"])
attention = DynamicScaledDotProductAttention(
max_seq_len=cfg.max_seq_len,
block_size=cfg.block_size,
config=cfg_hf,
device=device,
local_windows_len=cfg.local_windows_len,
topk=cfg.second_select_num,
threads_num=8,
anchor_type=cfg.anchor_type,
kv_type=cfg.kv_type,
dense_layer_num=cfg.dense_layer_num,
anchor_num=cfg.anchor_num,
block_selection_mode=cfg.head_select_mode,
layer_step=cfg.layer_step,
token_step=cfg.token_step,
preselect_block=cfg.preselect_block,
preselect_block_count=cfg.preselect_block_count,
prefill_chunk_size=cfg.chunk_size,
)
The constructor allocates GPU tensors for active blocks and CPU buffers for staging evicted blocks, establishing the memory hierarchy that enables 139K-token contexts on limited VRAM.
Step 4: Execute Chunked Prefill
Long sequences cannot process in a single forward pass due to memory constraints. Use chunk_prefill from utils.py to process the prompt incrementally, updating the block-wise cache with each segment.
from ktransformers.util.utils import chunk_prefill
def prefill_long_context(input_ids: torch.Tensor):
"""Process prompt in chunks to populate block-wise KV-cache."""
seq_len = input_ids.shape[1]
# long_context uses internal block cache, not standard StaticCache
past_key_values = None
for start in range(0, seq_len, cfg.chunk_size):
end = min(start + cfg.chunk_size, seq_len)
chunk = input_ids[:, start:end]
position_ids = torch.arange(start, end, device=device)
logits = chunk_prefill(
inputs=chunk,
cache_position=position_ids,
past_key_values=past_key_values
)
return logits # Final chunk logits for next-token prediction
Each chunk updates the global block cache; blocks exceeding local_windows_len are scored for importance and candidates for offloading to CPU or disk tiers.
Step 5: Generate Tokens with Block-Wise Decoding
After prefilling, generate subsequent tokens using decode_wrapper, which continuously manages block residency. The wrapper updates the preselect_block_table and executes the appropriate attention kernel based on block location.
from ktransformers.util.utils import decode_wrapper
def generate_long_context(
input_ids: torch.Tensor,
max_new_tokens: int = 256
):
"""Generate tokens using tiered KV-cache."""
# Prefill phase
_ = prefill_long_context(input_ids)
generated = input_ids.clone()
for _ in range(max_new_tokens):
position_ids = torch.tensor([generated.shape[1]], device=device)
next_token = decode_wrapper(
next_token=None,
position_ids=position_ids,
cache_position=position_ids,
cuda_graph_runner=None,
past_key_values=None, # Block cache maintained internally
inputs=generated,
seq_length=generated.shape[1],
)
generated = torch.cat([generated, next_token.unsqueeze(0)], dim=1)
if next_token.item() == tokenizer.eos_token_id:
break
return generated
The decode_wrapper automatically evicts low-importance blocks to CPU/Disk when GPU memory pressure exceeds thresholds, preserving context continuity across the full 139K-token window.
Optional: Enable Disk-Backed KV-Cache
For contexts exceeding combined GPU and CPU memory, configure disk spillover via the kvc2 section in your YAML. The InferenceContext in balance_serve/inference/model_runner.py handles serialization of evicted blocks.
kvc2:
gpu_only: false # Allow CPU/Disk spill
cpu_memory_size_GB: 64 # CPU cache reservation
disk_path: "/tmp/kvc2_cache" # Overflow block storage
When blocks are evicted from CPU, they serialize to disk_path and deserialize on-demand during attention computation, adding latency but preserving the full context window.
Complete Implementation Example
The following script demonstrates end-to-end 139K-token inference using the block-wise architecture:
# demo_139k_context.py
import torch
from ktransformers.server.config.config import Config
from ktransformers.util.utils import load_model, chunk_prefill, decode_wrapper
# Load configuration with long_context parameters
cfg = Config()
# Initialize model with block-wise attention
model, tokenizer = load_model(
cfg.model["name_or_path"],
device=cfg.model["device"],
mode="long_context",
)
# Prepare 139K-token synthetic prompt
prompt_text = "Detailed analysis of transformer scaling laws. " * 3500
input_ids = tokenizer(prompt_text, return_tensors="pt")["input_ids"]
if input_ids.shape[1] > cfg.max_seq_len:
input_ids = input_ids[:, :cfg.max_seq_len]
# Chunked prefill
print("Prefilling context...")
_ = chunk_prefill(
inputs=input_ids,
cache_position=torch.arange(input_ids.shape[1], device=cfg.model["device"]),
past_key_values=None,
)
# Generate completion
print("Generating...")
generated = input_ids.clone()
for i in range(100):
pos_id = torch.tensor([generated.shape[1]], device=cfg.model["device"])
next_tok = decode_wrapper(
next_token=None,
position_ids=pos_id,
cache_position=pos_id,
cuda_graph_runner=None,
past_key_values=None,
inputs=generated,
seq_length=generated.shape[1],
)
generated = torch.cat([generated, next_tok.unsqueeze(0)], dim=1)
if next_tok.item() == tokenizer.eos_token_id:
break
output = tokenizer.decode(generated[0], skip_special_tokens=True)
print(output)
On a single RTX 4090 (24GB VRAM), this configuration uses approximately 6GB GPU memory for the local window while maintaining the full 139K-token context through tiered storage.
Summary
- Configure the
long_contextsection in YAML withmax_seq_len: 139000and appropriateblock_size(typically 128). - Load the model using
mode="long_context"to activate the block-wise attention pathway inarchive/ktransformers/util/utils.py. - Prefill using
chunk_prefillto process long prompts incrementally without exhausting VRAM. - Generate via
decode_wrapper, which dynamically manages block residency across GPU, CPU, and Disk tiers. - Scale to 139K tokens by enabling disk-backed caching in the
kvc2configuration section when necessary.
Frequently Asked Questions
What is the maximum context length supported by KTransformers?
KTransformers supports context lengths up to 139,000 tokens or more, depending on your max_seq_len configuration and available tiered storage (CPU RAM and disk). The practical limit is determined by your willingness to trade latency for context length as blocks spill to slower storage tiers.
How does block-wise attention reduce memory usage?
Instead of storing the entire KV-cache for all 139K tokens in GPU memory, DynamicScaledDotProductAttention partitions the cache into blocks (default 128 tokens). Only the local_windows_len (e.g., 4096 tokens) and high-importance blocks remain on GPU; the framework offloads remaining blocks to CPU or disk according to the anchor_type selection policy.
Can I use this with multi-GPU setups?
Yes. While the example demonstrates single-GPU inference, you can distribute blocks across multiple GPUs by adjusting the device configuration. The head_select_mode parameter supports both shared and per-head block masks, allowing flexible sharding strategies across GPU clusters.
What happens when a block is evicted to disk?
When GPU and CPU memory are saturated, the InferenceContext in archive/ktransformers/server/balance_serve/inference/model_runner.py serializes evicted blocks to the disk_path specified in your YAML. During generation, if the model requires an off-disk block, the system deserializes it back to CPU, then potentially to GPU, introducing a latency penalty proportional to block size and disk speed.
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 →