How to Optimize Memory Usage When Running Llama 2 Inference with Large Batch Sizes
Reduce the max_batch_size parameter and implement chunked inference to prevent GPU out-of-memory errors when processing large workloads with the official Meta Llama 2 model.
The official meta-llama/llama repository stores substantial intermediate activation data in GPU memory during text generation, particularly within the attention mechanism's key-value cache. As batch sizes increase, these cached tensors grow linearly and often dominate total memory consumption. Understanding how to configure and partition these allocations is essential to optimize memory usage when running Llama 2 inference with large batch sizes without sacrificing generation quality.
Understanding Memory Bottlenecks in Llama 2
The KV-Cache Tensor Allocation
In llama/model.py, the Attention module pre-allocates persistent GPU buffers for keys and values during initialization. At lines 36-44, the code creates tensors with shape max_batch_size × max_seq_len × n_local_kv_heads × head_dim for both cache_k and cache_v. These allocations persist for the entire generation duration, meaning doubling your batch size doubles the cache memory footprint immediately.
Runtime Batch Size Guards
The generator enforces these limits strictly at runtime. In llama/generation.py lines 58-61, the Llama.generate method contains the assertion bsz <= params.max_batch_size, which hard-stops execution if you attempt to exceed the configured limit established during model construction.
Memory Optimization Strategies
Lower the Maximum Batch Size
The most direct approach reduces max_batch_size in ModelArgs (defined at line 30 of llama/model.py). Pass a smaller integer when building the model through Llama.build()—the model can process any batch smaller than or equal to this limit, while larger workloads must be split into chunks.
Implement Chunked Inference
Split logical batches into sub-batches that run sequentially. This keeps each iteration under the KV-cache limit while allowing you to process arbitrary total volumes of prompts. The chunk size must respect the runtime guard in Llama.generate that validates against params.max_batch_size.
Leverage Half-Precision Arithmetic
The codebase already defaults to torch.cuda.HalfTensor (line 18 of llama/generation.py), automatically halving memory consumption for weights and activations compared to float32. For newer GPUs such as the A100, switching to bfloat16 offers similar memory savings with better numerical stability for large models.
Eliminate Excessive Padding
The generator pads sequences to total_len (calculated as max_seq_len or max_gen_len + max_prompt_len). Reducing max_gen_len in your text_completion or chat_completion calls shrinks the working tensor sizes proportionally, lowering peak memory usage.
Advanced CPU Offloading
For extreme memory constraints, modify Attention.__init__ to allocate KV caches on CPU rather than GPU by replacing .cuda() with .to('cpu'), then copying needed slices to GPU on-demand during forward passes. This trades generation latency for substantial memory budget savings.
Implementation Examples
Building with a Reduced Batch Size
Configure the model builder with conservative memory limits to prevent out-of-memory errors during initialization.
from llama.generation import Llama
# Set a modest batch size to keep GPU memory low
llama = Llama.build(
ckpt_dir="checkpoints",
tokenizer_path="tokenizer.model",
max_seq_len=2048,
max_batch_size=8, # <-- smaller than the default 32
model_parallel_size=1,
seed=42,
)
The max_batch_size argument is forwarded to ModelArgs (see lines 13-15 of llama/generation.py), directly controlling the cache allocation size in the attention layers.
Chunked Inference for Large Workloads
Process arbitrarily large prompt lists by splitting them into chunks that respect the runtime batch limit.
def chunked_text_completion(
llama: Llama,
prompts: list[str],
chunk_size: int = 8, # must be ≤ llama.model.params.max_batch_size
**kwargs,
):
"""Run text completion on a large list of prompts by splitting into smaller chunks."""
results = []
for i in range(0, len(prompts), chunk_size):
chunk = prompts[i : i + chunk_size]
results.extend(llama.text_completion(chunk, **kwargs))
return results
# Example usage – 32 prompts with a model built for max_batch_size=8
answers = chunked_text_completion(
llama,
prompts=[f"Prompt {i}" for i in range(32)],
temperature=0.7,
top_p=0.9,
)
This function respects the runtime guard (assert bsz <= params.max_batch_size) by never submitting a batch larger than the configured limit.
Verifying KV Cache Dimensions
Inspect the actual allocated cache shapes to confirm your memory optimizations.
import torch
def print_kv_cache_sizes(llama: Llama):
# Access the first layer's Attention module to inspect cache shapes
att = llama.model.layers[0].attention
print("Key cache shape :", att.cache_k.shape)
print("Value cache shape :", att.cache_v.shape)
print_kv_cache_sizes(llama)
With max_batch_size=8 and max_seq_len=2048, the printed shapes will be torch.Size([8, 2048, n_local_kv_heads, head_dim]), confirming the reduced footprint compared to larger batch configurations.
Switching to BFloat16 on Supported Hardware
For newer GPUs that favor bfloat16 precision, explicitly set the default tensor type before model construction.
# Only needed on newer GPUs (e.g., A100) that favor bfloat16
torch.set_default_tensor_type(torch.cuda.BFloatTensor)
Make this change before calling Llama.build() so all parameters and KV-cache buffers use the new precision.
Summary
- KV-cache tensors in
Attention.__init__(lines 36-44 ofllama/model.py) consume memory proportional tomax_batch_size × max_seq_len, making them the primary bottleneck for large batches. - Reduce
max_batch_sizeinModelArgs(line 30 ofllama/model.py) to immediately lower cache allocations, then use chunked inference to process large workloads incrementally. - Half-precision defaults are already active in the repository (line 18 of
llama/generation.py), but you can upgrade tobfloat16on supported hardware for better stability. - Runtime assertions in
Llama.generate(lines 58-61 ofllama/generation.py) enforce batch limits, requiring chunking strategies for processing prompt lists that exceed your configured maximum.
Frequently Asked Questions
What consumes the most GPU memory during Llama 2 inference?
The KV-cache tensors dominate memory usage. According to the source code in llama/model.py, the Attention module allocates persistent GPU buffers for keys and values with shape max_batch_size × max_seq_len × n_local_kv_heads × head_dim at lines 36-44. These caches exist for the full duration of generation, scaling linearly with batch size and sequence length.
Can I process batches larger than max_batch_size without modifying the source code?
No, the runtime explicitly prevents this. In llama/generation.py lines 58-61, the Llama.generate method asserts bsz <= params.max_batch_size, raising an error if violated. Instead, implement chunked inference to split large logical batches into smaller sub-batches that each respect the configured limit.
Does Llama 2 automatically use mixed precision to save memory?
Yes. The repository sets torch.set_default_tensor_type(torch.cuda.HalfTensor) at line 18 of llama/generation.py, forcing all model weights and activations into float16 by default. Additionally, the @torch.inference_mode() decorator applied to Transformer.forward and Llama.generate disables gradient tracking, further reducing memory overhead during inference.
How can I verify my memory optimizations are working?
Inspect the actual cache tensor shapes at runtime. Access the attention module through llama.model.layers[0].attention and print cache_k.shape and cache_v.shape. These tensors should reflect your configured max_batch_size and max_seq_len values, confirming that reduced parameters translate directly into smaller GPU allocations.
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 →