Skip to main content
Create your own

Dynamic Batching for Efficient Inference

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:

  1. Schedule: The scheduler examines all waiting and running requests and composes an optimal batch for a single forward pass.
  2. Execute: The engine runs the model for that single batch.
  3. 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:

  1. Calling scheduler.schedule_next_batch() to get the plan.
  2. Gathering tokens and metadata from the scheduled sequences into batched tensors (input_ids, attention_mask, etc.).
  3. Executing the model() forward pass.
  4. Scattering the results (next token logits, updated KV cache) back to the individual Request objects managed by the scheduler.
Static vs. Continuous Batching for LLM Inference
A quick visual reminder of our goal. We are implementing the logic for the bottom panel (Continuous Batching), where the system dynamically fills GPU capacity by mixing and matching requests iteration by iteration.

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"
}

Can't find a good explanation? Sign up and we'll make it for you

Sign up