Inference Optimization Guide¶
This guide covers techniques for optimizing inference performance in the LLM framework.
KV Cache¶
KV Cache stores computed key and value tensors during autoregressive generation, avoiding redundant computation.
How It Works¶
sequenceDiagram
participant Model
participant Cache as KVCache
participant Attn as Attention
Note over Model,Attn: Initial prompt (tokens 1-10)
Model->>Cache: Store K, V for tokens 1-10
Note over Model,Attn: Generate token 11
Model->>Attn: Compute Q for token 11 only
Cache->>Attn: Provide K, V for tokens 1-10
Attn->>Attn: Attend to all 11 tokens
Attn->>Cache: Update with K, V for token 11
Without cache: O(n²) attention per token With cache: O(n) attention per token
Pre-Allocated KVCache¶
The KVCache class provides efficient, pre-allocated buffers:
from llm.core.kv_cache import KVCache
# Create cache for inference
cache = KVCache(
max_batch_size=1,
max_seq_len=512, # Maximum generation length
num_kv_heads=8, # From model config
head_dim=64, # hidden_size / num_heads
device="cuda",
dtype=torch.float16,
)
# For multi-layer models, create one per layer
caches = KVCache.from_model_config(
max_batch_size=1,
max_seq_len=512,
num_layers=12,
num_kv_heads=8,
head_dim=64,
device="cuda",
dtype=torch.float16,
)
Using with DecoderModel¶
from llm.core.kv_cache import KVCache, reset_all_caches
# Create caches
caches = KVCache.from_model_config(
max_batch_size=1,
max_seq_len=512,
num_layers=model.num_layers,
num_kv_heads=model.num_kv_heads,
head_dim=model.hidden_size // model.num_heads,
device=device,
dtype=model.dtype,
)
# Generation loop
input_ids = torch.tensor(tokenizer.encode("Hello, world!"), dtype=torch.long).unsqueeze(0) # [1, S]
for _ in range(max_new_tokens):
# Forward with cache — caches update in-place, and forward returns
# ``(logits, kv_caches)`` when ``use_cache=True``; unpack both.
logits, caches = model(input_ids, kv_caches=caches, use_cache=True)
# Get next token: [0, -1, :] = batch 0, last position
next_token = logits[0, -1, :].argmax(dim=-1).unsqueeze(0) # [1, 1]
input_ids = next_token # Only pass new token for next step
# Reset for new sequence
reset_all_caches(caches)
Memory Benefits¶
| Approach | Memory Pattern | Fragmentation |
|---|---|---|
torch.cat (legacy) |
Grows each step | High |
KVCache (pre-alloc) |
Fixed upfront | None |
Continuous Batching¶
For high-throughput serving, use the ContinuousBatchingEngine which supports iteration-level scheduling.
Engine Setup¶
from llm.serving.batch_engine import ContinuousBatchingEngine
from llm.models.decoder import DecoderModel
from llm.tokenization.simple_tokenizer import SimpleCharacterTokenizer
# Load model and tokenizer
model = DecoderModel(
vocab_size=32000,
hidden_size=512,
num_layers=6,
num_heads=8,
max_seq_len=512,
)
tokenizer = SimpleCharacterTokenizer(["H", "e", "l", "o", "W", "r", "d"]) # cover "Hello" / "World"
# Create engine (model and tokenizer required upfront)
engine = ContinuousBatchingEngine(
model=model,
tokenizer=tokenizer,
device="cuda",
max_batch_size=16,
max_seq_len=512,
)
Request Processing¶
from llm.serving.schemas import GenerationRequest
# Add requests
req1 = GenerationRequest(prompt="Hello", max_new_tokens=50)
req2 = GenerationRequest(prompt="World", max_new_tokens=50)
engine.add_request(req1)
engine.add_request(req2)
# Run inference steps
while engine.scheduler.has_pending_work:
engine.step() # One iteration handles all active sequences
Key Features¶
| Feature | Description |
|---|---|
| Iteration-level scheduling | Multiple requests processed per step |
| Slot-based KV cache | Pre-allocated memory pool |
| Mixed prefill/decode | New and ongoing sequences batched together |
| Automatic padding | Handles variable-length inputs |
Async Lifecycle (T2 #23)¶
Since the T2 #23 refactor, the engine exposes two complementary APIs:
engine.step()— synchronous, for scripts / tests.engine.step_async()— asynchronous wrapper for FastAPI handlers.
Both APIs decompose into three phases:
[lock] _lock_step_pre — slot alloc, prefix-cache lookup,
batch tensor construction
[free] _forward_and_sample — model forward + sampling (expensive)
[lock] _lock_step_post — append tokens, free slots, set status
The lock is released during the model forward so other workers
can pre-/post-compute in parallel. step_async offloads the forward
to a thread via asyncio.to_thread(...), letting the FastAPI
event loop keep serving health checks, /metrics scrapes, and other
in-flight requests while a forward pass runs.
# Async usage in a FastAPI handler
from fastapi.concurrency import run_in_threadpool
async def stream_one(prompt: str):
while True:
stats = await engine.step_async()
# ... yield decoded token to the client ...
The StepStats(scheduled, total_active_slots) returned by either
path feeds the llm_batch_fill_ratio gauge (see
Serving Metrics).
Serving Metrics (Prometheus)¶
prometheus-fastapi-instrumentator already emits generic HTTP RED
metrics (rate, errors, duration) per route. The serving tier also
publishes domain-specific metrics so operators can see what the model
is actually doing — not just that the route returned 200.
All metrics live in src/llm/serving/metrics.py and are exposed at
/metrics alongside the HTTP RED metrics.
Metrics reference¶
| Metric | Type | Labels | Source |
|---|---|---|---|
llm_tokens_generated_total |
Counter | endpoint |
observed per successful generation |
llm_tokens_per_request |
Histogram | endpoint |
distribution of completion tokens (16/64/256/1024/4096 buckets) |
llm_request_duration_seconds |
Histogram | endpoint, status |
end-to-end request duration (0.05/0.25/1/5/30 buckets) |
llm_batch_fill_ratio |
Gauge | — | ContinuousBatchingEngine.set_step_observer callback |
llm_kv_cache_hit_ratio |
Gauge | — | set by callers observing prefix-cache hits |
llm_inflight_requests |
Gauge | — | incremented while a request holds the semaphore |
Endpoints contributing to the endpoint label: generate,
batch_generate, chat_completions.
Example PromQL queries¶
# p95 latency per endpoint (seconds)
histogram_quantile(0.95,
sum by (le, endpoint) (rate(llm_request_duration_seconds_bucket[5m]))
)
# Throughput (tokens / second) by endpoint
sum by (endpoint) (rate(llm_tokens_generated_total[1m]))
# Batch utilization — fraction of slots in use over time
avg_over_time(llm_batch_fill_ratio[5m])
# Saturation signal — sustained near-100% fill with rising p95
# means the engine is throughput-bound.
llm_batch_fill_ratio > 0.8
and
histogram_quantile(0.95,
sum by (le) (rate(llm_request_duration_seconds_bucket[5m]))
) > 10
# KV-cache hit rate — fraction of prefix lookups served from cache
avg_over_time(llm_kv_cache_hit_ratio[10m])
Wiring a custom observer¶
The engine's set_step_observer(callback) hook fires once per
engine.step() under the step lock, with the latest StepStats. Use
it to publish gauges (e.g. slot utilization) or to drive
adaptive batching decisions:
from llm.serving.batch_engine import ContinuousBatchingEngine
from llm.serving.metrics import METRICS
engine = ContinuousBatchingEngine.from_serving_config(config, model, tokenizer)
engine.set_step_observer(METRICS.record_batch_fill_ratio)
Pass None to clear a previously installed observer.
Startup configuration log line¶
On lifespan startup, the server emits one structured JSON line tagged
event: server_config. This is the canonical record of what is actually
running — useful for incident triage when the question is "which model
is on this box?" or "is prefix cache actually on?".
The line is emitted once via _log_server_config in src/llm/serving/api.py,
keyed on event="server_config" so it is greppable and Prometheus-style
log shippers can route on it without parsing free-form text.
Example (CPU dummy model, no API key):
{"event": "server_config", "model_class": "DecoderModel",
"param_count_total": 113125, "param_count_trainable": 113125,
"dtype": "torch.float32", "device": "cpu",
"max_seq_len": 128, "attn_impl": "mha", "mlp_impl": "mlp",
"generation_backend": "eager", "enable_prefix_cache": false,
"use_paged_attention": false, "api_key_set": false}
Fields:
| Field | Notes |
|---|---|
model_class |
type(model).__name__ — DecoderModel for built-in, custom for third-party |
param_count_total |
every parameter (trainable + frozen) |
param_count_trainable |
requires_grad=True only — useful sanity check that PEFT froze the base |
dtype / device |
from the first parameter; "unknown" if the model has none |
max_seq_len |
from ServingConfig.max_seq_len |
attn_impl / mlp_impl |
registry keys (mha / mlp for built-ins; custom values for plugins) |
generation_backend |
eager / batched — which backend loop is in use(speculative 是 Python API 专用后端,需要 target + draft 双模型,不能经 serving 配置) |
enable_prefix_cache |
bool |
use_paged_attention |
bool |
api_key_set |
bool ONLY — the key value itself is never logged |
Operators normally grep for this line:
# Tail server logs and pull only the startup config line
uv run llm-serve 2>&1 | tee /tmp/llm-serve.log | grep '"event": "server_config"'
## Server Binding & Authentication Policy
The server **refuses to start** when it would bind to a non-loopback
address (`0.0.0.0`, `*`, LAN IPs, public hostnames) **without** an
`api_key` configured. This is a fail-closed guard: binding to a public
interface with no authentication exposes the inference endpoint to
anonymous network access, which is almost never intended.
### Loopback allow-list
| Host value | Loopback? | Safe without `api_key`? |
| ----------------- | --------- | ----------------------- |
| `127.0.0.1` | Yes | Yes |
| `127.*.*.*` | Yes | Yes |
| `::1` | Yes | Yes |
| `localhost` | Yes | Yes |
| `0.0.0.0` | No | No — requires `api_key` |
| `192.168.1.100` | No | No — requires `api_key` |
### How to run on a public interface
Two options, both required:
```bash
# Option A: keep loopback (default, safe for local development)
LLM_SERVING_HOST=127.0.0.1 uv run llm-serve
# Option B: bind to all interfaces + set an API key (containerized deployment)
LLM_SERVING_HOST=0.0.0.0 \
LLM_SERVING_API_KEY=$(openssl rand -hex 32) \
uv run llm-serve
When api_key is set, all /generate, /batch_generate,
/v1/chat/completions, and /metrics requests require a matching
Authorization: Bearer <key> header or X-API-Key: <key> header.
Requests without a valid key receive 403 Unauthorized. /health is
the only route that stays public.
Re-extract from a saved log file¶
grep '"event": "server_config"' /var/log/llm-serve.log | tail -1 | jq .
The `api_key_set` field intentionally carries the **presence** of an
API key, never its value — see `src/llm/serving/auth.py:52` for the
`hmac.compare_digest` timing-safe comparison.
## Grouped Query Attention (GQA)
GQA reduces KV cache memory by sharing KV heads across multiple query heads.
```text
MHA: Q=32, K=32, V=32 (32 KV pairs)
GQA: Q=32, K=8, V=8 (8 KV pairs, 4x memory reduction)
MQA: Q=32, K=1, V=1 (1 KV pair, 32x memory reduction)
Configuration¶
model = DecoderModel(
hidden_size=1024,
num_heads=16, # Query heads
num_kv_heads=4, # KV heads (GQA: 4:1 ratio)
)
Sliding Window Attention¶
Limits attention to recent tokens only, reducing memory for very long sequences.
model = DecoderModel(
hidden_size=512,
num_heads=8,
window_size=256, # Only attend to last 256 tokens
)
Trade-offs¶
| Window Size | Memory | Long-range Recall |
|---|---|---|
| 128 | Very low | Limited |
| 512 | Low | Good |
| 2048 | Medium | Excellent |
| None | High | Full |
Flash Attention 2 (opt-in)¶
For Ampere / Hopper hardware, the project exposes an explicit Flash
Attention 2 backend through ATTENTION_REGISTRY as attn_impl="flash_attn".
It is opt-in because flash-attn ships CUDA wheels only and
should not be a hard runtime dependency.
Install¶
# uv
uv sync --extra perf
# pip
pip install 'llm[perf]'
# or, on a CUDA host:
pip install flash-attn
uv.lock pins flash-attn as an sdist, so a plain uv sync --extra perf
compiles it from source on first install — this needs a CUDA toolkit
(nvcc) and can take a long time. On a CUDA host you can instead pull
pre-built GPU wheels from the
Astral GPU indexes, which publish flash-attn
builds for each supported CUDA + PyTorch combination:
# Pick the CUDA index that matches your torch build (e.g. cu126 / cu128 /
# cu130). uv pins flash-attn to this index with `explicit = true`, so no
# other dependency is affected.
uv add flash-attn --index astral-cu126=https://wheels.astral.sh/simple/cu126/
Astral wheels encode the CUDA + PyTorch pair in the local version tag
(e.g. 2.8.3.post1+cu.12.6.torch.2.12). Make sure the wheel you select
matches your installed torch CUDA build, PyTorch version, Python version,
and platform — a mismatch is the most common cause of ImportError:
flash_attn ... no module named at runtime. If no pre-built wheel matches
your exact environment (e.g. a very new torch/Python combo), fall back to
the source build above or pip install flash-attn on a CUDA host.
torch 2.13 / Python 3.14 note: as of this writing the Astral indexes only publish
flash-attnwheels up totorch 2.12, so ontorch 2.13+ Python 3.14 you must build from source. The build requires annvccwhose major.minor matchestorch.version.cuda(e.g. CUDA 13.x fortorch 2.13.0+cu130) — torch'scpp_extensionrefuses a mismatched compiler even though CUDA runtime is forward-compatible. Do not useFLASH_ATTENTION_SKIP_CUDA_BUILD=TRUE(the default in this repo's[tool.uv.extra-build-variables]): it produces an importable package with no kernels and every flash test fails deep inforward.
Configuration¶
from llm.models.decoder import DecoderModel
model = DecoderModel(
vocab_size=32000,
hidden_size=1024,
num_layers=12,
num_heads=16,
attn_impl="flash_attn", # uses FlashAttention when CUDA is available
)
If flash-attn is not installed, the registry entry still resolves
but constructing the module raises a clear ImportError pointing to
the install command. CPU-only hosts should keep the default
attn_impl="mha".
Trade-offs vs. mha¶
| Aspect | mha (default) |
flash_attn |
|---|---|---|
| Backend | torch.nn.functional.scaled_dot_product_attention |
flash_attn.flash_attn_func |
| Custom attn_mask | Supported | Not supported — falls back via the engine layer if you need padding masks |
| Sliding window | Supported (PyTorch SDPA) | Supported via window_size= on dense seqs; + padding mask still needs flash_attn_varlen_func |
| Hardware | CPU + CUDA + MPS | CUDA only |
| Long-context perf | O(S²) memory peaks | Streaming softmax, O(S) memory |
For training or long-context decode on supported hardware, prefer
flash_attn. For variable-length sequences with padding masks, stick
with mha until flash_attn_varlen_func lands (see Tier 3 follow-up).
Publishing models to HuggingFace Hub¶
The reverse of compat/hf_loader.from_pretrained is in
compat/hf_publisher.py. Trained DecoderModels can be saved
locally in HF-compatible format and pushed to a Hub repo, so the
model is loadable by both this project and HF's transformers.
Install¶
# uv
uv sync --extra compat
# pip
pip install 'llm[compat]'
# or, directly:
pip install safetensors huggingface_hub
Save locally¶
from llm.compat.hf_publisher import save_pretrained
from llm.models.decoder import DecoderModel
model = DecoderModel(
vocab_size=32000,
hidden_size=4096,
num_layers=32,
num_heads=32,
attn_impl="mha",
mlp_impl="mlp",
use_glu=True,
)
# ... train model ...
save_pretrained(model, save_directory="./my-llm")
# → writes config.json + model.safetensors (Llama-shaped)
Push to Hub¶
from llm.compat.hf_publisher import push_to_hub
push_to_hub(
model,
repo_id="alice/my-llm",
private=False,
commit_message="Initial upload",
)
# Returns "https://huggingface.co/alice/my-llm"
Roundtrip guarantee¶
save_pretrained writes a directory that the existing
from_pretrained can load back into an equivalent model. The
q/k/v split + concat, plus fc1 ↔ up_proj / fc2 ↔ down_proj
translation, is exercised in
tests/compat/test_hf_publisher.py::test_save_pretrained_roundtrip_through_from_pretrained.
Limitations (Tier 3 follow-ups)¶
- Only
architecture="llama"is wired through. Qwen/Mistral architectures are recognized bydetect_architecturebut the publisher raisesValueError. Full support lands with the FSDP-e2e + multi-arch work (Tier 3 #2). - Tokenizer artifacts are not yet auto-saved alongside the model.
Bundle them with
tokenizer.save_pretrained(save_directory)before pushing, or open a follow-up.
Speculative decoding¶
For long-context decode on large target models, speculative decoding
(Leviathan et al., 2023) can give 2–3× throughput. A small draft
model speculates gamma candidate tokens ahead of the large
target model; the target scores all candidates in a single
forward pass and either accepts (probabilistically, preserving the
target distribution) or samples a correction token.
Configuration¶
from llm.generation.backends import SpeculativeDecodingBackend
backend = SpeculativeDecodingBackend(
target_model=target, # the "expensive" model
draft_model=draft, # smaller model, same vocab
gamma=5, # speculative tokens per round
)
# Or via the registry:
from llm.generation.registry import get_generation_backend
backend = get_generation_backend(
"speculative",
target_model=target,
draft_model=draft,
gamma=5,
)
When it helps¶
| Scenario | Speedup? |
|---|---|
| Long-context decode (≥ 256 tokens), draft well-aligned with target | ✓ 2–3× typical |
| Short prompts (< 32 tokens) | ✗ overhead dominates |
| Draft very different from target (e.g., mixed-model families) | ✗ acceptance rate collapses |
| Greedy decoding | Marginal — most tokens already accepted in eager mode |
Sampling¶
Speculative decoding is distribution-preserving: sampling from
SpeculativeDecodingBackend with the same temperature / top_k /
top_p as the eager backend produces text from the same
distribution as the target model alone. Greedy
(temperature=0.0) turns the acceptance test into plain argmax
equality — a draft token is accepted iff its argmax equals the
target's argmax at that position, and the correction on a mismatch is
the target's own argmax. A well-aligned draft therefore wins gamma
tokens plus one bonus per round; a misaligned draft is still cut at
the first mismatch (see "acceptance rate collapses" in the table
above) — greedy never forces the two argmaxes to agree.
Limitations¶
- Both models must share vocabulary. The draft is not auto-trained — pick a smaller sibling or distill the target.
- The current slice runs the eager algorithm (each round
rebuilds the context tensor). Wiring it into
ContinuousBatchingEnginefor serving-tier concurrency is a Tier 3 follow-up.
Inference Checklist¶
- ✅ Use KVCache for autoregressive generation
- ✅ Enable GQA if model supports it (check num_kv_heads)
- ✅ Consider sliding window for very long sequences
- ✅ Use appropriate dtype (fp16/bf16 for GPU)
- ✅ Consider Flash Attention 2 on Ampere/Hopper (
attn_impl="flash_attn") - ✅ Merge LoRA weights before inference (
merge_lora())
Performance Comparison¶
| Technique | Latency | Memory | Quality |
|---|---|---|---|
| Baseline | 1.0x | 1.0x | 100% |
| + KVCache | 0.3x | ~1.0x | 100% |
| + GQA (4:1) | 0.25x | 0.25x | ~99% |
| + Sliding Window | 0.2x | 0.15x | ~95%* |
*Quality depends on task; long-range dependencies may suffer.