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).

Downloads last month

-

Downloads are not tracked for this model. How to track
Video Preview
loading