Contents
CartPole is a great first project, but let's be honest — nobody gets genuinely excited watching a stick balance. Snake is different. Watching an agent go from suicidally crashing into walls to weaving cleanly around its own tail chasing apples is the RL project that actually feels like you built something cool, not just a proof of concept.
I built this exact project right after CartPole clicked for me, specifically because Snake forces you to solve a problem CartPole conveniently sidesteps: designing a state representation from scratch. CartPole hands you four clean numbers. Snake doesn't hand you anything — you have to decide what the agent even gets to "see."
By the end of this tutorial, you'll have a Snake-playing agent trained with DQN, and you'll understand exactly why state design matters more than almost anything else in this project. IMO, this is where RL stops feeling like a tutorial and starts feeling like actual engineering :)
Why Snake Is a Genuinely Different Challenge Than CartPole
CartPole's state space came pre-packaged: four numbers, done. Snake gives you raw pixel data or grid positions, and you have to engineer a state representation that actually captures what matters — otherwise your agent trains forever and learns nothing useful.
The action space is small (4 directions, or 3 if you exclude reversing into yourself), which is genuinely convenient. The state space is the real challenge — deciding what information the agent needs to make good decisions, without drowning it in irrelevant detail. The reward structure needs real thought — sparse rewards (only rewarding eating food) train painfully slowly compared to a well-shaped reward signal.
Ever wondered why some Snake AI tutorials train in minutes while others take hours for mediocre results? It's almost always the state representation and reward shaping, not the neural network architecture.
If you're coming from CartPole, our CartPole tutorial covers the DQN fundamentals — replay buffer, epsilon-greedy exploration, and the learning loop — that this project builds directly on top of.
Figure 1: Training a DQN agent to play Snake — from random crashes to intelligent food-chasing behavior
Setting Up the Environment
We'll build the game with Pygame and train using PyTorch — this combination shows up consistently across the strongest community implementations of this exact project.
python -m venv venv
source venv/bin/activate # Windows: venv\Scripts\activate
pip install pygame torch numpy matplotlib
Matplotlib here isn't optional decoration — plotting your training score over time is genuinely the fastest way to tell whether your agent is actually improving or just flailing productively.
Building the Snake Game Environment
Before any learning happens, we need a working game the agent can actually interact with. Here's a simplified grid-based version.
import pygame
import random
import numpy as np
from collections import namedtuple
Point = namedtuple('Point', 'x, y')
class SnakeGame:
def __init__(self, w=640, h=480, block_size=20):
self.w = w
self.h = h
self.block_size = block_size
self.reset()
def reset(self):
self.direction = "RIGHT"
self.head = Point(self.w // 2, self.h // 2)
self.snake = [self.head,
Point(self.head.x - self.block_size, self.head.y),
Point(self.head.x - 2 * self.block_size, self.head.y)]
self.score = 0
self.food = None
self._place_food()
self.frame_iteration = 0
def _place_food(self):
x = random.randint(0, (self.w - self.block_size) // self.block_size) * self.block_size
y = random.randint(0, (self.h - self.block_size) // self.block_size) * self.block_size
self.food = Point(x, y)
if self.food in self.snake:
self._place_food()
That recursive _place_food call handles the edge case where food spawns directly on the snake's body — a small detail that's easy to skip and annoying to debug later when your agent inexplicably eats food that was never actually there.
Designing the State Representation (The Part That Actually Matters)
This is genuinely the most important design decision in the whole project. A commonly used, well-tested representation boils the entire game down to 11 boolean values:
3 values for danger detection — is there danger straight ahead, to the right, or to the left of the snake's current direction? 4 values for current direction — one-hot encoded as up, down, left, right. 4 values for food location relative to the head — is food up, down, left, or right of the snake?
def get_state(game):
head = game.snake[0]
point_l = Point(head.x - 20, head.y)
point_r = Point(head.x + 20, head.y)
point_u = Point(head.x, head.y - 20)
point_d = Point(head.x, head.y + 20)
dir_l = game.direction == "LEFT"
dir_r = game.direction == "RIGHT"
dir_u = game.direction == "UP"
dir_d = game.direction == "DOWN"
state = [
(dir_r and game.is_collision(point_r)) or
(dir_l and game.is_collision(point_l)) or
(dir_u and game.is_collision(point_u)) or
(dir_d and game.is_collision(point_d)),
(dir_u and game.is_collision(point_r)) or
(dir_d and game.is_collision(point_l)) or
(dir_l and game.is_collision(point_u)) or
(dir_r and game.is_collision(point_d)),
(dir_d and game.is_collision(point_r)) or
(dir_u and game.is_collision(point_l)) or
(dir_r and game.is_collision(point_u)) or
(dir_l and game.is_collision(point_d)),
dir_l, dir_r, dir_u, dir_d,
game.food.x < head.x, game.food.x > head.x,
game.food.y < head.y, game.food.y > head.y
]
return np.array(state, dtype=int)
This 11-value representation is deliberately relative to the snake's own orientation, not absolute screen coordinates. That choice matters enormously — an absolute-coordinate state means the agent has to relearn "danger is close" for every possible position on the board, while a relative state generalizes immediately regardless of where the snake happens to be.
Why Not Just Feed It Raw Pixels?
Fair question. You genuinely could train a convolutional network directly on pixel data, and that's how more advanced versions of this project work. But raw pixels require dramatically more training time and compute to extract the same information this handcrafted 11-value state gives you for free. Start with the engineered state. Pixel-based CNN input is a legitimate "advanced version" project once this one clicks.
For a deeper understanding of how state design affects RL performance, our evaluating reinforcement learning algorithms guide covers the metrics and methods used to assess whether your agent is actually learning useful behavior.
Designing the Reward Structure
Sparse rewards — only rewarding the agent when it eats food — technically work, but they train painfully slowly, since the agent gets almost no useful signal for most of its actions.
REWARD_EAT = 10
REWARD_DEATH = -10
REWARD_STEP = 0
Notice this is deliberately simple: +10 for eating food, -10 for dying, 0 otherwise. No per-step penalty here, though some implementations add a small negative reward per step to discourage the agent from stalling indefinitely without making progress.
Adding a small negative reward for wasted time is worth experimenting with once your basic version trains successfully — it pushes the agent toward more efficient food-seeking behavior instead of just avoiding death indefinitely.
The DQN Network Architecture
Compared to CartPole, the network itself barely changes — this problem's complexity lives in the state design, not the model architecture.
import torch
import torch.nn as nn
class DQN(nn.Module):
def __init__(self, input_size=11, hidden_size=256, output_size=3):
super().__init__()
self.net = nn.Sequential(
nn.Linear(input_size, hidden_size),
nn.ReLU(),
nn.Linear(hidden_size, output_size)
)
def forward(self, x):
return self.net(x)
Notice the output size is 3, not 4 — straight, right turn, or left turn, relative to the snake's current direction, rather than absolute up/down/left/right. This prevents the agent from ever needing to "consider" reversing directly into itself, since that action simply doesn't exist in this framing.
The Training Loop
The core training mechanics mirror CartPole almost exactly — replay buffer, epsilon-greedy exploration, periodic learning steps.
import random
from collections import deque
class Agent:
def __init__(self):
self.n_games = 0
self.epsilon = 0
self.gamma = 0.9
self.memory = deque(maxlen=100_000)
self.model = DQN()
self.optimizer = torch.optim.Adam(self.model.parameters(), lr=0.001)
def get_action(self, state):
self.epsilon = 80 - self.n_games
final_move = [0, 0, 0]
if random.randint(0, 200) < self.epsilon:
move = random.randint(0, 2)
else:
state_tensor = torch.tensor(state, dtype=torch.float)
prediction = self.model(state_tensor)
move = torch.argmax(prediction).item()
final_move[move] = 1
return final_move
Notice this epsilon decay approach is slightly different from CartPole's exponential decay — it decreases linearly with the number of games played, tied directly to training progress rather than a fixed multiplicative rate. Both approaches work; this one just ties exploration more directly to how experienced the agent already is.
Storing and Learning From Experience
def remember(self, state, action, reward, next_state, done):
self.memory.append((state, action, reward, next_state, done))
def train_long_memory(self):
if len(self.memory) > 1000:
mini_sample = random.sample(self.memory, 1000)
else:
mini_sample = self.memory
states, actions, rewards, next_states, dones = zip(*mini_sample)
self.train_step(states, actions, rewards, next_states, dones)
This is genuinely the same replay-buffer pattern from CartPole — randomly sampling past experiences to break harmful correlation between consecutive moves. If that trick worked for balancing a pole, it works just as well for not running into your own tail.
Running the Full Training Loop
def train():
scores = []
record = 0
agent = Agent()
game = SnakeGame()
while True:
state_old = get_state(game)
final_move = agent.get_action(state_old)
reward, done, score = game.play_step(final_move)
state_new = get_state(game)
agent.remember(state_old, final_move, reward, state_new, done)
if done:
game.reset()
agent.n_games += 1
agent.train_long_memory()
if score > record:
record = score
print(f"Game {agent.n_games}, Score {score}, Record {record}")
scores.append(score)
Run this, and expect a genuinely recognizable pattern: the first 50-100 games look almost identical to random flailing, scores mostly under 5. Somewhere past that point, scores start climbing steadily, and by a few hundred games, an agent using this exact state design commonly reaches average scores in the 20s-40s range — genuinely competent play for a game this simple.
Common Mistakes When Building This Yourself
I've hit most of these myself, so treat this as a shortcut past the frustration.
Using absolute screen coordinates instead of relative state. This forces the agent to relearn "danger nearby" separately for every board position, dramatically slowing training. Forgetting to prevent the snake from reversing into itself. If your action space includes "move opposite direction," you'll get instant, confusing deaths that have nothing to do with the agent's actual decision quality. Setting the reward for eating too low relative to death penalty. If dying costs -10 but eating only gives +1, the agent learns overly cautious behavior instead of actively pursuing food. Not tracking a rolling average score. Individual game scores bounce around a lot; plotting the raw scores makes training look far noisier than it actually is — track a moving average instead.
Visualizing Training Progress
import matplotlib.pyplot as plt
def plot_scores(scores):
plt.clf()
plt.plot(scores)
plt.xlabel("Game Number")
plt.ylabel("Score")
plt.savefig("training_progress.png")
Don't skip this. Watching the raw print statements scroll by makes it genuinely hard to judge whether training is actually working. A plotted score trend makes the improvement (or lack of it) immediately obvious in a way scrolling numbers never will.
Where to Go From Here
Once your agent reliably scores in the 20s or better, a few natural extensions make sense.
Try feeding raw pixel data through a convolutional network instead of the handcrafted 11-value state — genuinely harder, but a good next step once you understand why the engineered state worked in the first place. Experiment with reward shaping — try adding a small reward for moving closer to food, not just for eating it, and observe how that changes learned behavior. Swap DQN for PPO using Stable-Baselines3 to see how a more modern algorithm handles the exact same problem.
If you want to formalize your RL knowledge after this project, our best courses for RL in robotics and game AI covers Hugging Face's Deep RL Course, NVIDIA's Physical AI path, and Stanford CS234.
Wrapping This Up
Building a Snake AI with DQN uses the exact same core mechanics as CartPole — replay buffer, epsilon-greedy exploration, periodic learning steps — but forces you to actually think about state design and reward shaping instead of getting them handed to you. That relative, 11-value state representation is doing most of the heavy lifting here, more than any tweak to the network architecture would.
Remember that sparse rewards train slower than shaped ones, and a relative state representation generalizes dramatically better than absolute screen coordinates. FYI, once this version trains successfully, swapping in raw pixels and a CNN feels like a natural next challenge rather than starting over from scratch :)
Now go tweak the reward values and watch what kind of snake behavior emerges. A slightly wrong reward for eating versus surviving will teach you more about reward shaping than any explanation ever could.