SLM-Tetris-Arena / tetris.py
DedeProGames's picture
Track the last action (piece, rotation, column, lines, options)
1a3bdb6 verified
Raw History Blame Contribute Delete
8.86 kB
"""Deterministic Tetris engine used by the LM Tetris Arena.
The engine works at the *placement* level: for each new piece we enumerate every
(rotation, column) that can be hard-dropped straight down, simulate the result and
describe the outcome in plain English. Players (language models) only choose among
those outcomes, so every move they make is legal.
"""
from __future__ import annotations
import random
from dataclasses import dataclass, field
WIDTH = 10
HEIGHT = 20
PIECE_NAMES = ["I", "O", "T", "S", "Z", "J", "L"]
# (row, col) cells of each piece in its spawn orientation
_BASE = {
"I": [(0, 0), (0, 1), (0, 2), (0, 3)],
"O": [(0, 0), (0, 1), (1, 0), (1, 1)],
"T": [(0, 1), (1, 0), (1, 1), (1, 2)],
"S": [(0, 1), (0, 2), (1, 0), (1, 1)],
"Z": [(0, 0), (0, 1), (1, 1), (1, 2)],
"J": [(0, 0), (1, 0), (1, 1), (1, 2)],
"L": [(0, 2), (1, 0), (1, 1), (1, 2)],
}
COLOR_ID = {name: i + 1 for i, name in enumerate(PIECE_NAMES)} # 1..7
LINE_POINTS = {0: 0, 1: 100, 2: 300, 3: 500, 4: 800}
def _normalize(cells):
mr = min(r for r, _ in cells)
mc = min(c for _, c in cells)
return tuple(sorted((r - mr, c - mc) for r, c in cells))
def _rotations(cells):
out, seen = [], set()
cur = list(cells)
for _ in range(4):
n = _normalize(cur)
if n not in seen:
seen.add(n)
out.append(n)
cur = [(c, -r) for r, c in cur] # rotate 90 degrees clockwise
return out
ROTATIONS = {name: _rotations(cells) for name, cells in _BASE.items()}
class PieceSequence:
"""7-bag randomizer. Two sequences built from the same seed are identical."""
def __init__(self, seed: int):
self._rng = random.Random(seed)
self._pieces: list[str] = []
def __getitem__(self, i: int) -> str:
while len(self._pieces) <= i:
bag = PIECE_NAMES[:]
self._rng.shuffle(bag)
self._pieces.extend(bag)
return self._pieces[i]
# ----------------------------------------------------------------------------
# Board helpers
# ----------------------------------------------------------------------------
def empty_grid():
return [[0] * WIDTH for _ in range(HEIGHT)]
def column_heights(grid):
heights = [0] * WIDTH
for c in range(WIDTH):
for r in range(HEIGHT):
if grid[r][c]:
heights[c] = HEIGHT - r
break
return heights
def count_holes(grid):
holes = 0
for c in range(WIDTH):
seen_block = False
for r in range(HEIGHT):
if grid[r][c]:
seen_block = True
elif seen_block:
holes += 1
return holes
def bumpiness(heights):
return sum(abs(heights[i] - heights[i + 1]) for i in range(WIDTH - 1))
# ----------------------------------------------------------------------------
# Natural-language description of an outcome
# ----------------------------------------------------------------------------
_LINES_TXT = {
0: "clears no lines",
1: "clears one line",
2: "clears two lines",
3: "clears three lines",
4: "clears four lines at once",
}
def describe(lines: int, new_holes: int, max_height: int, bump: int, landing: int) -> str:
"""Neutral, factual English description of what a move does.
Values are bucketed into words on purpose: tiny LMs handle words far better
than numbers, and a small closed set of phrases keeps scoring cacheable.
`landing` = how many rows above the lowest column the piece comes to rest.
"""
if landing <= 0:
landing_txt = "drops the piece into the lowest part of the board"
elif landing <= 2:
landing_txt = "drops the piece a little above the lowest part of the board"
else:
landing_txt = "drops the piece far above the lowest part of the board"
lines_txt = _LINES_TXT[lines]
if new_holes < 0:
holes_txt = "removes some holes"
elif new_holes == 0:
holes_txt = "creates no new holes"
elif new_holes == 1:
holes_txt = "creates one new hole"
else:
holes_txt = "creates several new holes"
if max_height <= 4:
height_txt = "keeps the stack very low"
elif max_height <= 8:
height_txt = "keeps the stack low"
elif max_height <= 12:
height_txt = "makes the stack high"
elif max_height <= 16:
height_txt = "makes the stack very high"
else:
height_txt = "brings the stack close to the top"
if bump <= 4:
surface_txt = "leaves the surface flat"
elif bump <= 10:
surface_txt = "leaves the surface a little uneven"
else:
surface_txt = "leaves the surface very bumpy"
return f"{landing_txt}, {lines_txt}, {holes_txt}, {height_txt} and {surface_txt}"
@dataclass
class Candidate:
rotation: int
x: int
cells: list # absolute (row, col) cells where the piece lands
grid: list # board after placement and line clears
lines: int
holes: int
new_holes: int
max_height: int
agg_height: int
bump: int
landing: int
description: str
def enumerate_placements(grid, piece: str, holes_before: int | None = None) -> list[Candidate]:
if holes_before is None:
holes_before = count_holes(grid)
heights = column_heights(grid)
tops = [HEIGHT - h for h in heights] # first filled row index (HEIGHT if empty)
lowest = min(heights)
color = COLOR_ID[piece]
out = []
for rot_idx, shape in enumerate(ROTATIONS[piece]):
w = max(c for _, c in shape) + 1
bottom = {}
for r, c in shape:
bottom[c] = max(bottom.get(c, -1), r)
for x in range(WIDTH - w + 1):
# landing offset: the piece stops when any column touches the stack
y = min(tops[x + pc] - 1 - br for pc, br in bottom.items())
cells = [(r + y, c + x) for r, c in shape]
if any(r < 0 for r, _ in cells):
continue # would stick out of the top: illegal (top-out)
g = [row[:] for row in grid]
for r, c in cells:
g[r][c] = color
kept = [row for row in g if not all(row)]
lines = HEIGHT - len(kept)
if lines:
g = [[0] * WIDTH for _ in range(lines)] + kept
hs = column_heights(g)
holes = count_holes(g)
bump = bumpiness(hs)
mh = max(hs)
landing = (HEIGHT - 1 - max(r for r, _ in cells)) - lowest
out.append(
Candidate(
rotation=rot_idx,
x=x,
cells=cells,
grid=g,
lines=lines,
holes=holes,
new_holes=holes - holes_before,
max_height=mh,
agg_height=sum(hs),
bump=bump,
landing=landing,
description=describe(lines, holes - holes_before, mh, bump, landing),
)
)
return out
@dataclass
class TetrisGame:
seed: int
grid: list = field(default_factory=empty_grid)
score: int = 0
lines: int = 0
pieces: int = 0
tetrises: int = 0
alive: bool = True
last_cells: list = field(default_factory=list)
last_description: str = ""
last_value: float | None = None
last_piece: str = ""
last_rotation: int = 0
last_rotations: int = 1
last_col: int = 0
last_lines: int = 0
last_options: int = 0
def __post_init__(self):
self.sequence = PieceSequence(self.seed)
self._holes = 0
@property
def current_piece(self) -> str:
return self.sequence[self.pieces]
@property
def next_piece(self) -> str:
return self.sequence[self.pieces + 1]
def candidates(self) -> list[Candidate]:
return enumerate_placements(self.grid, self.current_piece, self._holes)
def apply(self, cand: Candidate, value: float | None = None, options: int = 0):
piece = self.current_piece
self.last_piece = piece
self.last_rotation = cand.rotation
self.last_rotations = len(ROTATIONS[piece])
self.last_col = cand.x + 1 # leftmost column of the piece, 1..10
self.last_lines = cand.lines
self.last_options = options
self.grid = cand.grid
self._holes = cand.holes
self.lines += cand.lines
self.score += LINE_POINTS[cand.lines]
self.tetrises += int(cand.lines == 4)
self.pieces += 1
# cells of the placed piece that survived line clears (for highlighting)
self.last_cells = cand.cells if cand.lines == 0 else []
self.last_description = cand.description
self.last_value = value
def top_out(self):
self.alive = False