Introduction
In our last lesson, you designed a StallFreeScheduler that produces a hybrid batch of prefill chunks and decode tokens for each iteration. This scheduler creates the plan for what the GPU should work on next to balance throughput and latency. However, a plan is only useful if it can be executed.
This lesson bridges the gap between planning and execution. Your goal is to integrate the scheduler into a full inference loop. You will build the core of an inference engine that takes the scheduler's plan, dynamically assembles the necessary tensors for the model, executes a single forward pass, and processes the results. This is the heart of a modern LLM serving system, where the scheduler and the execution engine work in a tight, continuous cycle.
1. The Scheduler-Engine Interaction Model
The system we're building operates on a principle called iteration-level scheduling. Instead of the scheduler preparing a large, static batch that runs to completion, it interacts with the execution engine on a per-iteration basis.
This creates a highly dynamic and responsive system. Let's watch a short video that illustrates this concept.
OSDI '22 - Orca: A Distributed Serving System for Transformer-Based Generative Models
The OSDI '22 presentation on the Orca serving system provides a clear, concise explanation of iteration-level scheduling.
Watch the segment from 06:56 to 08:47. Pay close attention to how the scheduler selects requests for just one iteration, sends them to the engine, and then immediately re-evaluates the pool of requests (including newly arrived ones) for the next iteration. This is the loop we are about to implement.
As the video shows, the core logic is a loop:
- Schedule: The scheduler examines all waiting and running requests and composes an optimal batch for a single forward pass.
- Execute: The engine runs the model for that single batch.
- Update: The engine processes the results, updates the state of the requests (e.g., adding a new token, marking a request as complete), and the loop repeats.
This constant communication allows the engine to add new requests as soon as capacity becomes available, maximizing GPU utilization.
To get a more detailed mental model of this process, let's see how it's realized in a production system like vLLM.
How the VLLM inference engine works?
The video 'How the VLLM inference engine works?' gives a fantastic, step-by-step visualization of the entire process, connecting the scheduler's decisions to the underlying data structures.
Watch the detailed walkthrough from 48:44 to 1:02:40. Focus on the interplay between the 'waiting queue' and the 'running queue'. Observe how the scheduler forms a continuous batch (49:50), how blocks are allocated from the free list (52:20), and most importantly, how the scheduler moves requests to the 'running queue' after their prefill is done (59:00). Notice how a completed request's resources are immediately freed and added back to the pool (1:01:40). This is the dynamic process your code will orchestrate.
2. The Core Logic: The Inference Loop
Now, let's translate this conceptual model into a concrete implementation structure. The inference loop is the central component that orchestrates the work.
At a high level, the loop's logic can be described with the following pseudocode, which is inspired by the architecture of systems like Hugging Face's Text Generation Inference (TGI).
LLM Inference at scale with TGI - Hugging Face
The Hugging Face blog post on TGI provides pseudocode that clearly illustrates the continuous batching algorithm. Let's examine it to formalize our loop structure.
Read the pseudocode in the section 'The Router: Queueing and Continuous Batching'. Note the main while batch: loop. Inside, it adds new requests if budget allows, decodes the current batch, and filters out completed requests. This is the cycle of 'add, decode, filter' that you will implement.
Our implementation will follow a similar pattern. We'll create an InferenceEngine class that contains our StallFreeScheduler and the main inference loop. The loop will be responsible for:
- Calling
scheduler.schedule_next_batch()to get the plan. - Gathering tokens and metadata from the scheduled sequences into batched tensors (
input_ids,attention_mask, etc.). - Executing the
model()forward pass. - Scattering the results (next token logits, updated KV cache) back to the individual
Requestobjects managed by the scheduler.

3. Implementation Exercise: The Inference Engine
It's time to write the code. Below is the skeleton for an InferenceEngine. It includes the Request and ScheduledSequence dataclasses from our previous lesson, along with your completed StallFreeScheduler.
Your task is to complete the run method. This method will house the main inference loop. To keep you focused on the high-level orchestration, the complex tensor manipulation for gathering inputs and scattering outputs will be handled by provided helper methods: _prepare_model_inputs and _process_model_outputs. Your job is to call them in the correct sequence within the loop.
import torch
from collections import deque
from dataclasses import dataclass, field
from typing import List, Deque, Dict, Optional
# --- Data Structures (from previous lesson) ---
@dataclass
class Request:
id: int
prompt_tokens: List[int]
max_new_tokens: int
tokens_processed: int = 0
output_tokens: List[int] = field(default_factory=list)
# In a real system, we'd use an EOS token ID.
is_finished: bool = False
def is_prefill_done(self) -> bool:
return self.tokens_processed >= len(self.prompt_tokens)
def is_complete(self) -> bool:
return self.is_finished or (self.is_prefill_done() and len(self.output_tokens) >= self.max_new_tokens)
@dataclass
class ScheduledSequence:
request_id: int
tokens: List[int]
is_decode: bool
# We add position_ids_start to handle heterogeneous batches
position_ids_start: int
# --- Scheduler (from previous lesson, provided for you) ---
class StallFreeScheduler:
def __init__(self, token_budget: int, max_batch_size: int):
self.token_budget = token_budget
self.max_batch_size = max_batch_size
self.waiting_queue: Deque[Request] = deque()
self.active_pool: Dict[int, Request] = {}
def add_request(self, request: Request):
self.waiting_queue.append(request)
def get_active_and_waiting_count(self):
return len(self.active_pool), len(self.waiting_queue)
def schedule_next_batch(self) -> List[ScheduledSequence]:
next_batch: List[ScheduledSequence] = []
current_tokens = 0
# 1. Evict completed requests
completed_ids = [req_id for req_id, req in self.active_pool.items() if req.is_complete()]
for req_id in completed_ids:
# In a real engine, we'd also free the KV cache here
del self.active_pool[req_id]
# 2. Prioritize ongoing decodes
for req_id, req in self.active_pool.items():
if req.is_prefill_done() and not req.is_complete():
if current_tokens + 1 <= self.token_budget and len(next_batch) < self.max_batch_size:
next_batch.append(ScheduledSequence(request_id=req_id, tokens=[req.output_tokens[-1]], is_decode=True, position_ids_start=len(req.prompt_tokens) + len(req.output_tokens) -1))
current_tokens += 1
# 3. Handle partially completed prefills
for req_id, req in self.active_pool.items():
if not req.is_prefill_done():
remaining_prompt = len(req.prompt_tokens) - req.tokens_processed
chunk_size = min(remaining_prompt, self.token_budget - current_tokens)
if chunk_size > 0 and len(next_batch) < self.max_batch_size:
start = req.tokens_processed
end = req.tokens_processed + chunk_size
next_batch.append(ScheduledSequence(request_id=req_id, tokens=req.prompt_tokens[start:end], is_decode=False, position_ids_start=start))
current_tokens += chunk_size
req.tokens_processed += chunk_size
break
# 4. Admit new requests
while self.waiting_queue and len(self.active_pool) < self.max_batch_size and current_tokens < self.token_budget:
req = self.waiting_queue[0]
if len(self.active_pool) + 1 > self.max_batch_size: break
req = self.waiting_queue.popleft()
self.active_pool[req.id] = req
chunk_size = min(len(req.prompt_tokens), self.token_budget - current_tokens)
if chunk_size > 0:
next_batch.append(ScheduledSequence(request_id=req.id, tokens=req.prompt_tokens[:chunk_size], is_decode=False, position_ids_start=0))
current_tokens += chunk_size
req.tokens_processed += chunk_size
else:
self.waiting_queue.appendleft(req)
del self.active_pool[req.id]
break
return next_batch
# --- Your Task: The Inference Engine ---
class InferenceEngine:
def __init__(self, model, tokenizer, scheduler: StallFreeScheduler, device='cuda'):
self.model = model.to(device)
self.tokenizer = tokenizer
self.scheduler = scheduler
self.device = device
# This will store the KV cache for active requests
self.kv_cache: Dict[int, Optional[torch.Tensor]] = {}
def _prepare_model_inputs(self, scheduled_batch: List[ScheduledSequence]) -> Dict[str, torch.Tensor]:
"""Gathers tokens, position_ids, and KV caches into a single batch."""
input_tokens = []
position_ids = []
# This mask helps us find where each sequence starts in the concatenated tensor
cu_seqlens = [0]
# We need to separate prefill and decode caches
past_key_values_prefill = []
past_key_values_decode = []
decode_indices = [] # Indices of decode sequences in the batch
for i, seq in enumerate(scheduled_batch):
input_tokens.extend(seq.tokens)
cu_seqlens.append(cu_seqlens[-1] + len(seq.tokens))
pos = torch.arange(seq.position_ids_start, seq.position_ids_start + len(seq.tokens), dtype=torch.long)
position_ids.append(pos)
# KV cache handling
if seq.is_decode:
decode_indices.append(i)
if seq.request_id in self.kv_cache and self.kv_cache[seq.request_id] is not None:
past_key_values_decode.append(self.kv_cache[seq.request_id])
else: # Prefill
# Prefill uses KV cache only if it's a continuation of a chunked prefill
if seq.request_id in self.kv_cache and self.kv_cache[seq.request_id] is not None:
past_key_values_prefill.append(self.kv_cache[seq.request_id])
# In a real system like vLLM with PagedAttention, this would be much more complex.
# For simplicity, we can only batch requests that are all starting (no KV cache)
# or are all decoding (all have KV cache). A mix is complex to handle without PagedAttention.
# Our scheduler logic primarily creates batches of decodes + ONE prefill, so we can simplify.
return {
"input_ids": torch.tensor(input_tokens, dtype=torch.long, device=self.device).unsqueeze(0),
"position_ids": torch.cat(position_ids).to(self.device).unsqueeze(0),
# In a simplified HuggingFace model, we might not need a complex attention mask if position_ids are correct.
# For a real implementation, you'd build a block-diagonal mask.
}
def _process_model_outputs(self, outputs, scheduled_batch: List[ScheduledSequence]):
"""Scatters the model outputs back to the individual requests."""
# We only care about the logit for the *last* token of each sequence in the batch
# For decode, it's 1 token. For prefill, it's the last token of the chunk.
# Simplified next token selection (greedy)
next_token_ids = torch.argmax(outputs.logits[:, -1, :], dim=-1)
next_token_id = next_token_ids[0].item()
# For this simplified engine, we assume the output corresponds to the LAST sequence added to the batch
# A real engine would need to map logits back to each sequence correctly.
last_seq = scheduled_batch[-1]
req_id = last_seq.request_id
req = self.scheduler.active_pool.get(req_id)
if req:
req.output_tokens.append(next_token_id)
# Check for a simplified stop condition
if next_token_id == self.tokenizer.eos_token_id:
req.is_finished = True
# Update KV Cache
# This is also simplified. A real engine would update/append to the cache for each request.
# self.kv_cache[req_id] = outputs.past_key_values
def run(self):
"""
The main inference loop. It continuously schedules and executes batches
until all requests are processed.
"""
iteration = 0
# Continue as long as there are requests active or waiting
while self.scheduler.get_active_and_waiting_count()[0] > 0 or self.scheduler.get_active_and_waiting_count()[1] > 0:
iteration += 1
print(f"\n--- Iteration {iteration} ---")
# --- YOUR IMPLEMENTATION GOES HERE ---
# 1. Get the next batch to process from the scheduler.
scheduled_batch = self.scheduler.schedule_next_batch()
# 2. Check if the scheduler returned any work. If not, the GPU is idle for this cycle.
# Break the loop if there's no more work to be done.
if not scheduled_batch:
print("Scheduler returned no batch. Ending run.")
break
print(f"Scheduled {len(scheduled_batch)} sequences.")
for seq in scheduled_batch:
print(f" - Req ID {seq.request_id}: {'Decode' if seq.is_decode else 'Prefill'} {len(seq.tokens)} tokens.")
# 3. Prepare the dictionary of inputs for the model by calling the helper method.
# (The helper is simplified and only processes one sequence at a time for this exercise)
if len(scheduled_batch) > 0:
# To simplify this exercise, we'll process only the last scheduled item
# which is most likely to be a new prefill chunk. A real engine handles the full batch.
single_sequence_batch = [scheduled_batch[-1]]
model_inputs = self._prepare_model_inputs(single_sequence_batch)
# 4. Execute a single forward pass of the model with the prepared inputs.
with torch.no_grad():
outputs = self.model(**model_inputs)
# 5. Process the model's outputs using the helper method to update the request state.
self._process_model_outputs(outputs, single_sequence_batch)
```grasp
{
"type": "exercise",
"id": "53566bee-8e5d-47ec-ae08-0595861976ff"
}
# Log the output for the request we just processed
req_id = single_sequence_batch[0].request_id
if req_id in self.scheduler.active_pool:
req = self.scheduler.active_pool[req_id]
current_output = self.tokenizer.decode(req.output_tokens)
print(f" Request {req_id} current output: '{current_output}'")
print("\n--- Run Complete ---")
Example Usage
if name == 'main':
from transformers import AutoModelForCausalLM, AutoTokenizer
# This is a mock model for demonstration purposes.
# In a real scenario, you'd load a real model.
class MockModel(torch.nn.Module):
def __init__(self, vocab_size=50257):
super().__init__()
self.vocab_size = vocab_size
self.dummy_param = torch.nn.Parameter(torch.empty(0))
def forward(self, input_ids, **kwargs):
# Simple mock logic: "generates" the next token ID
next_token_logit = torch.randn(1, 1, self.vocab_size, device=self.dummy_param.device)
# Make token 500 (' test') likely
next_token_logit[:,:,-1] = 10
from collections import namedtuple
ModelOutput = namedtuple("ModelOutput", ["logits", "past_key_values"])
return ModelOutput(logits=next_token_logit, past_key_values=None)
# Setup
# model_name = "gpt2"
# tokenizer = AutoTokenizer.from_pretrained(model_name)
# model = AutoModelForCausalLM.from_pretrained(model_name)
# if tokenizer.pad_token is None:
# tokenizer.pad_token = tokenizer.eos_token
tokenizer = AutoTokenizer.from_pretrained("gpt2")
tokenizer.pad_token = tokenizer.eos_token
model = MockModel(vocab_size=len(tokenizer))
scheduler = StallFreeScheduler(token_budget=1024, max_batch_size=8)
engine = InferenceEngine(model, tokenizer, scheduler)
# Add some requests
prompts = [
"The capital of France is",
"The best programming language is",
"Large language models are"
]
for i, prompt in enumerate(prompts):
prompt_tokens = tokenizer.encode(prompt)
engine.scheduler.add_request(Request(id=i, prompt_tokens=prompt_tokens, max_new_tokens=10))
# Run the engine
engine.run()
The provided code includes a simplified `_prepare_model_inputs` and `_process_model_outputs` to allow you to run this example. In a production system, these helpers would be significantly more complex to handle true parallel processing of a heterogeneous batch, which is the domain of technologies like PagedAttention. The key takeaway for you is the orchestration within the `run` loop.
### Conclusion
Congratulations, you have now implemented the central loop of a continuous batching inference engine. By integrating the scheduler's planning with the model's execution, you've created a system that can dynamically manage and process multiple requests to maximize efficiency.
**Key Takeaways:**
* The inference engine operates in a tight **`schedule -> execute -> update` loop**, driven by the scheduler.
* Integrating the scheduler requires a **"gather"** step, where inputs from different scheduled sequences are collated into a single batch for the model.
* It also requires a **"scatter"** step, where the batched output from the model is distributed back to update the state of individual requests.
* The complexity of handling these heterogeneous batches is why specialized components like vLLM's PagedAttention are so critical for performance, as they manage the underlying memory and tensor operations efficiently.
**Preview of the Next Lesson:**
Our current scheduler makes decisions based on a token budget and batch size, but it's blind to the most critical constraint: **KV cache memory**. It might schedule a new request only for the engine to discover there isn't enough VRAM to store its KV cache, leading to an out-of-memory error. In the next lesson, you will **design the data structures and interface for communication between the scheduler and a paged KV cache manager**, making the scheduler memory-aware and our engine far more robust.
```grasp
{
"type": "exercise",
"id": "db0155e1-2c6c-4920-afb9-10faa54b02ce"
}