snake-rl-ppo (ONNX)
A small convolutional network trained with PPO to play Snake on a 24 × 16 board. 5 MB, fp32 ONNX. No planner, no search: it reads the board and picks a move. It averages ~237 points (max 378) and plays in milliseconds, even on a Raspberry Pi 4.
It is the "RL" mode of laya-pi and the teacher of the Laya fine-tune trinacratech/snake-rl-room-onnx.
| file | what |
|---|---|
snake-rl.onnx |
input obs float32 (N, 6, 16, 24); outputs logits (N, 4) in order UP, DOWN, LEFT, RIGHT, and value (N) |
Observation
Six 16 × 24 planes (rows × columns, row 0 at the top). With body listed head first and
left[y, x] = len(body) - i for the i-th body cell (moves until that cell frees up):
| channel | content |
|---|---|
| 0 | occupied: left > 0 |
| 1 | left / (24 * 16) |
| 2 | left / len(body) |
| 3 | food: 1.0 on the food cell |
| 4 | fill: len(body) / (24 * 16) everywhere |
| 5 | hunger: min(moves_since_last_food / (24 * 16), 2.0) everywhere |
Moves into a wall, the body, or straight back are masked (logit set to −inf) before the argmax, as in training. Everything else the net decides itself; it can still box itself in.
import numpy as np, onnxruntime as ort
W, H = 24, 16
def observe(body, food, hunger):
n, cap = len(body), W * H
left = np.zeros((H, W), np.float32)
for i, (x, y) in enumerate(body): # body: [(x, y), ...] head first
left[y, x] = n - i
obs = np.zeros((1, 6, H, W), np.float32)
obs[0, 0] = left > 0
obs[0, 1] = left / cap
obs[0, 2] = left / n
if food:
obs[0, 3, food[1], food[0]] = 1.0
obs[0, 4] = n / cap
obs[0, 5] = min(hunger / cap, 2.0)
return obs
sess = ort.InferenceSession("snake-rl.onnx", providers=["CPUExecutionProvider"])
logits, value = sess.run(None, {"obs": observe(body, food, hunger)})
logits = np.where(legal_mask, logits[0], -np.inf) # legal_mask: 4 bools, UP DOWN LEFT RIGHT
move = ["UP", "DOWN", "LEFT", "RIGHT"][int(np.argmax(logits))]
laya-pi's snakeweb/rl_policy.py
is a complete player.
Training
- PPO, about 80 million moves of self-play in 3 hours, reward +1 for food and −1 for dying.
- Trained on a GPU copy of the game engine, checked move by move against the real engine: 300 games, 79,000 moves identical.
Evaluation
20 games, seeds 50000–50019, real engine. Maximum score 378 (the snake starts at length 6 on 384 cells).
| player | mean score | moves per food |
|---|---|---|
| this net (PyTorch) | 236.8 | ~25 |
| this net (ONNX fp32) | 226.4 (greedy ties diverge from PyTorch) | |
| Hamiltonian cycle + shortcuts | 378 (always wins) | 89 |
| Laya snake-rl-room, taught by this net | 102.2 |
On a Raspberry Pi 4 (Cortex-A72, ONNX Runtime fp32): ~18 ms per move with 3 threads; 10 games averaged 249 (min 93, max 336). On a desktop CPU: ~1.3 ms per move.
The net is fast and greedy (about 25 moves per food against the cycle's 90). It dies by boxing itself in: it never learned the cycle's patience.
Limitations
24 × 16 boards only. For demonstration, teaching, and as a teacher for distillation.
Credits
Snake rules from laya-mlx (Apache-2.0).