chessmamba-play / app.py
multimodalart's picture
multimodalart HF Staff
Enable show_error so runtime exceptions surface in the UI
ed51014 verified
Raw
History Blame Contribute Delete
21.3 kB
"""Play chess against ChessMamba (TobiasLogic/chessmamba).
ChessMamba is a ~16.8M-parameter selective state-space (Mamba / S6) chess engine
trained from scratch: it reads a game as a *sequence of moves*, keeps a compact
recurrent hidden state, and exposes three heads off that state — policy
(4096-way from-square x to-square), promotion (5-way) and value (scalar, tanh).
On top of that sits a policy-guided negamax with alpha-beta + quiescence.
This Space wires the author's own `model.py` / `search.py` / `chess_io.py`
(vendored verbatim from the model repo) into an interactive board, and surfaces
the raw policy / value heads so you can see what the state-space model actually
thinks before the search runs.
"""
from __future__ import annotations
import os
import re
import tempfile
import time
import chess
import chess.pgn
import chess.svg
import cairosvg
import gradio as gr
from huggingface_hub import hf_hub_download
from search import Searcher, DEFAULT_MAX_QDEPTH, DEFAULT_QS_TOP_K
MODEL_ID = "TobiasLogic/chessmamba"
# --------------------------------------------------------------------------- #
# Model loading (once, at module scope). ChessMamba is a tiny 16.8M-param model
# whose search is a Python-driven alpha-beta over batch-1 recurrent steps — it is
# latency-bound on tiny ops, so CPU is the right (and the author's own) device:
# `search.py` defaults to device="cpu" and torch.set_num_threads(1).
# --------------------------------------------------------------------------- #
CKPT_PATH = hf_hub_download(MODEL_ID, "ckpt/model.pt")
_BASE = Searcher(checkpoint_path=CKPT_PATH, device="cpu")
MODEL = _BASE.model
N_PARAMS = sum(p.numel() for p in MODEL.parameters())
print(f"[chessmamba] loaded {N_PARAMS:,} params from {CKPT_PATH}", flush=True)
def make_searcher(top_k: int = 10, max_depth: int = 6) -> Searcher:
"""A fresh Searcher per request (its root_node cache is mutable state),
sharing the single read-only model instance."""
s = Searcher(
checkpoint_path=None,
top_k=int(top_k),
max_depth=int(max_depth),
max_qdepth=DEFAULT_MAX_QDEPTH,
qs_top_k=DEFAULT_QS_TOP_K,
device="cpu",
)
s.model = MODEL
return s
# --------------------------------------------------------------------------- #
# Board rendering + click -> square mapping
# --------------------------------------------------------------------------- #
BOARD_PX = 520
# python-chess SVG layout: 8 * SQUARE_SIZE(45) + 2 * MARGIN(20) = 400 units.
_UNITS = 400.0
_SCALE = BOARD_PX / _UNITS
MARGIN_PX = 20 * _SCALE
SQUARE_PX = 45 * _SCALE
SEL_COLOR = "#f6b26b99"
DEST_COLOR = "#3d85c655"
def render_board(
board: chess.Board,
flip: bool = False,
selected: str = "",
arrow: chess.Move | None = None,
) -> str:
"""Render `board` to a PNG on disk and return its path."""
fill: dict[int, str] = {}
if selected:
try:
sq = chess.parse_square(selected)
except ValueError:
sq = None
if sq is not None and board.piece_at(sq) is not None:
fill[sq] = SEL_COLOR
for mv in board.legal_moves:
if mv.from_square == sq:
fill.setdefault(mv.to_square, DEST_COLOR)
arrows = []
if arrow is not None:
arrows.append(chess.svg.Arrow(arrow.from_square, arrow.to_square, color="#1f6feb99"))
svg = chess.svg.board(
board=board,
orientation=chess.BLACK if flip else chess.WHITE,
lastmove=board.peek() if board.move_stack else None,
check=board.king(board.turn) if board.is_check() else None,
fill=fill,
arrows=arrows,
size=BOARD_PX,
coordinates=True,
).encode("utf-8")
png = cairosvg.svg2png(bytestring=svg, output_width=BOARD_PX, output_height=BOARD_PX)
f = tempfile.NamedTemporaryFile(suffix=".png", delete=False)
f.write(png)
f.close()
return f.name
def xy_to_square(x: float, y: float, flip: bool) -> str:
"""Map a click on the rendered board image to a square name ('' if off-board)."""
fx = (x - MARGIN_PX) / SQUARE_PX
fy = (y - MARGIN_PX) / SQUARE_PX
if not (0 <= fx < 8 and 0 <= fy < 8):
return ""
col, row = int(fx), int(fy)
file_index = 7 - col if flip else col
rank_index = row if flip else 7 - row
return chess.square_name(chess.square(file_index, rank_index))
# --------------------------------------------------------------------------- #
# Game-state helpers. All state lives in (hidden) Gradio components so every
# handler is a pure function of its inputs.
# --------------------------------------------------------------------------- #
def board_from_history(history: str) -> chess.Board:
board = chess.Board()
for uci in (history or "").split():
try:
board.push_uci(uci)
except Exception:
break
return board
def history_of(board: chess.Board) -> str:
return " ".join(m.uci() for m in board.move_stack)
def movetext(board: chess.Board) -> str:
if not board.move_stack:
return "*(no moves yet)*"
game = chess.pgn.Game.from_board(board)
exporter = chess.pgn.StringExporter(headers=False, variations=False, comments=False)
text = game.accept(exporter)
return re.sub(r"\s*(1-0|0-1|1/2-1/2|\*)\s*$", "", text).strip()
def outcome_text(board: chess.Board) -> str:
if board.is_checkmate():
return f"♚ **Checkmate — {'White' if board.turn == chess.BLACK else 'Black'} wins.**"
if board.is_stalemate():
return "½ **Stalemate — draw.**"
if board.is_insufficient_material():
return "½ **Draw — insufficient material.**"
if board.is_seventyfive_moves() or board.is_fivefold_repetition():
return "½ **Draw.**"
return ""
def turn_text(board: chess.Board, human_white: bool) -> str:
over = outcome_text(board)
if over:
return over
check = " — **check!**" if board.is_check() else ""
human_turn = board.turn == (chess.WHITE if human_white else chess.BLACK)
return ("Your move." if human_turn else "ChessMamba is thinking…") + check
def legal_sans(board: chess.Board) -> list[str]:
return sorted(board.san(m) for m in board.legal_moves)
# --------------------------------------------------------------------------- #
# Reading the model's heads directly (policy priors + value)
# --------------------------------------------------------------------------- #
def analysis_md(board: chess.Board, played: str | None, elapsed: float | None,
top_k: int = 10) -> str:
"""Markdown panel showing the raw S6 policy/value heads at `board`."""
if board.is_game_over():
return f"### Engine\n\n{outcome_text(board)}"
s = make_searcher(top_k=top_k)
s.sync(board)
priors, value = s._policy_value(board, s.root_node)
if not priors:
return "### Engine\n\n*(no candidate moves — the position is drawn or over)*"
ranked = sorted(priors.items(), key=lambda kv: -kv[1])[: max(3, min(int(top_k), 8))]
mover = "White" if board.turn == chess.WHITE else "Black"
rows = ["| move | policy prior |", "| --- | --- |"]
for mv, p in ranked:
rows.append(f"| `{board.san(mv)}` | {p * 100:5.1f}% |")
head = "### ChessMamba's read of this position\n\n"
head += f"**Value head:** `{value:+.3f}` (+1 = winning for {mover}, −1 = losing)\n\n"
if played is not None:
timing = f" in {elapsed:.1f}s" if elapsed is not None else ""
head = (
f"### ChessMamba played `{played}`{timing}\n\n"
f"**Value head (now, {mover} to move):** `{value:+.3f}`\n\n"
)
head += "**Policy head — top candidate replies:**\n\n"
return head + "\n".join(rows)
# --------------------------------------------------------------------------- #
# Engine move
# --------------------------------------------------------------------------- #
def engine_move(board: chess.Board, movetime: float, top_k: int, max_depth: int):
"""Run the repo's policy-guided negamax and return (move, seconds)."""
s = make_searcher(top_k=top_k, max_depth=max_depth)
t0 = time.perf_counter()
mv = s.choose_move(board, ["go", "movetime", str(int(float(movetime) * 1000))])
return mv, time.perf_counter() - t0
# --------------------------------------------------------------------------- #
# Handlers
# --------------------------------------------------------------------------- #
PROMO_MAP = {"Queen": chess.QUEEN, "Rook": chess.ROOK, "Bishop": chess.BISHOP,
"Knight": chess.KNIGHT}
def _pack(board: chess.Board, human_white: bool, selected: str = "",
analysis: str = "", arrow: chess.Move | None = None, note: str = ""):
status = note or turn_text(board, human_white)
return (
render_board(board, flip=not human_white, selected=selected, arrow=arrow),
history_of(board),
selected,
status,
movetext(board),
analysis,
gr.update(choices=legal_sans(board), value=None),
)
def _apply_engine(board: chess.Board, human_white: bool, movetime: float,
top_k: int, max_depth: int):
mv, dt = engine_move(board, movetime, top_k, max_depth)
if mv is None:
return _pack(board, human_white, analysis=analysis_md(board, None, None, top_k))
san = board.san(mv)
board.push(mv)
return _pack(
board, human_white,
analysis=analysis_md(board, san, dt, top_k),
arrow=mv,
)
def _parse_move(board: chess.Board, text: str) -> chess.Move | None:
text = (text or "").strip()
if not text:
return None
for parser in (board.parse_san, board.parse_uci):
try:
mv = parser(text)
if mv in board.legal_moves:
return mv
except Exception:
continue
return None
def new_game(side: str, movetime: float, top_k: int, max_depth: int):
"""Start a fresh game; if the human plays Black, ChessMamba opens as White."""
human_white = side == "Play as White"
board = chess.Board()
if human_white:
return _pack(board, True, analysis=analysis_md(board, None, None, top_k))
return _apply_engine(board, False, movetime, top_k, max_depth)
def play_move(move_text: str, dropdown: str, history: str, side: str,
movetime: float, top_k: int, max_depth: int, promo: str):
"""Play the human's move (SAN or UCI), then let ChessMamba reply.
Args:
move_text: the human move in SAN (`Nf3`, `e4`, `O-O`) or UCI (`g1f3`).
dropdown: a legal move picked from the dropdown (used if move_text is empty).
history: space-separated UCI moves of the game so far.
side: "Play as White" or "Play as Black".
movetime: seconds of search time ChessMamba gets per move.
top_k: how many policy-ranked moves the search explores per node.
max_depth: maximum negamax depth.
promo: piece to promote to.
"""
human_white = side == "Play as White"
board = board_from_history(history)
if board.is_game_over():
yield _pack(board, human_white, analysis=analysis_md(board, None, None, top_k))
return
chosen = move_text if (move_text or "").strip() else dropdown
mv = _parse_move(board, chosen)
if mv is None:
hint = ", ".join(legal_sans(board)[:5])
yield _pack(board, human_white,
note=f"⚠️ `{chosen or '∅'}` is not a legal move here. Try e.g. {hint} …")
return
board.push(mv)
yield _pack(board, human_white, note="ChessMamba is thinking… 🐍")
if board.is_game_over():
yield _pack(board, human_white, analysis=analysis_md(board, None, None, top_k))
return
yield _apply_engine(board, human_white, movetime, top_k, max_depth)
def board_click(history: str, selected: str, side: str, movetime: float,
top_k: int, max_depth: int, promo: str, evt: gr.SelectData):
"""Click a piece, then click its destination square."""
human_white = side == "Play as White"
board = board_from_history(history)
x, y = float(evt.index[0]), float(evt.index[1])
sq_name = xy_to_square(x, y, flip=not human_white)
if board.is_game_over() or not sq_name:
yield _pack(board, human_white, selected="")
return
human_turn = board.turn == (chess.WHITE if human_white else chess.BLACK)
if not human_turn:
yield _pack(board, human_white, selected="")
return
sq = chess.parse_square(sq_name)
if not selected:
piece = board.piece_at(sq)
if piece is None or piece.color != board.turn:
yield _pack(board, human_white, selected="", note="Pick one of your own pieces.")
return
yield _pack(board, human_white, selected=sq_name,
note=f"`{sq_name}` selected — now click a destination.")
return
if sq_name == selected:
yield _pack(board, human_white, selected="")
return
from_sq = chess.parse_square(selected)
mv = chess.Move(from_sq, sq)
if mv not in board.legal_moves:
promoted = chess.Move(from_sq, sq, promotion=PROMO_MAP.get(promo, chess.QUEEN))
if promoted in board.legal_moves:
mv = promoted
else:
piece = board.piece_at(sq)
if piece is not None and piece.color == board.turn:
yield _pack(board, human_white, selected=sq_name,
note=f"`{sq_name}` selected — now click a destination.")
else:
yield _pack(board, human_white, selected="",
note=f"`{selected}{sq_name}` is not legal. Pick a piece again.")
return
board.push(mv)
yield _pack(board, human_white, note="ChessMamba is thinking… 🐍")
if board.is_game_over():
yield _pack(board, human_white, analysis=analysis_md(board, None, None, top_k))
return
yield _apply_engine(board, human_white, movetime, top_k, max_depth)
def undo(history: str, side: str, top_k: int):
"""Take back the last full move (yours + ChessMamba's)."""
human_white = side == "Play as White"
board = board_from_history(history)
for _ in range(2):
if board.move_stack:
board.pop()
return _pack(board, human_white, analysis=analysis_md(board, None, None, top_k))
def play_opening(first_move: str, movetime: float = 3.0, top_k: int = 10,
max_depth: int = 6):
"""Start a new game as White with `first_move`, then show ChessMamba's reply.
Args:
first_move: the opening move in SAN (`e4`, `d4`, `Nf3`, `c4`).
movetime: seconds of search time ChessMamba gets.
top_k: how many policy-ranked moves the search explores per node.
max_depth: maximum negamax depth.
"""
board = chess.Board()
mv = _parse_move(board, first_move)
if mv is None:
return _pack(board, True, analysis=analysis_md(board, None, None, top_k))[:6] + \
("Play as White",)
board.push(mv)
return _apply_engine(board, True, movetime, top_k, max_depth)[:6] + ("Play as White",)
# --------------------------------------------------------------------------- #
# UI
# --------------------------------------------------------------------------- #
INTRO = f"""
# 🐍♟️ Play ChessMamba
[`{MODEL_ID}`](https://huggingface.co/{MODEL_ID}) is a **{N_PARAMS/1e6:.1f}M-parameter selective
state-space model (Mamba / S6)** trained from scratch to play chess — no alpha-beta eval
function, no NNUE, no CNN+MCTS, and no Stockfish anywhere in its training. It reads the game
as a *sequence of moves*, advancing one compact recurrent hidden state per ply, and reads
**policy**, **promotion** and **value** heads straight off that state. A policy-guided
negamax with alpha-beta + quiescence sits on top.
**Click a piece, then click where it should go** (or type a move). ChessMamba replies, and the
panel on the right shows the raw policy/value heads behind its choice.
"""
CSS = """
#col-container { max-width: 1200px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
#board-img img { image-rendering: auto; }
"""
_START = chess.Board()
with gr.Blocks(title="Play ChessMamba") as demo:
with gr.Column(elem_id="col-container"):
gr.Markdown(INTRO)
history = gr.Textbox(value="", visible=False)
selected = gr.Textbox(value="", visible=False)
with gr.Row():
with gr.Column(scale=3):
board_img = gr.Image(
value=render_board(_START),
type="filepath",
label="Board — click a piece, then its destination",
elem_id="board-img",
interactive=False,
buttons=[],
height=BOARD_PX,
)
status_md = gr.Markdown("Your move.")
with gr.Column(scale=2):
with gr.Row():
move_in = gr.Textbox(
show_label=False,
placeholder="…or type a move: Nf3, e4, O-O, g1f3",
container=False,
scale=4,
)
play_btn = gr.Button("Play", variant="primary", scale=1)
legal_dd = gr.Dropdown(
label="…or pick a legal move",
choices=legal_sans(_START),
value=None,
)
analysis_out = gr.Markdown(
analysis_md(_START, None, None, 10), label="Engine"
)
moves_md = gr.Markdown("*(no moves yet)*", label="Moves")
with gr.Row():
new_btn = gr.Button("New game", variant="secondary")
undo_btn = gr.Button("Undo", variant="secondary")
with gr.Accordion("Engine settings", open=False):
with gr.Row():
side_radio = gr.Radio(
["Play as White", "Play as Black"],
value="Play as White",
label="Your side",
)
promo_radio = gr.Radio(
["Queen", "Rook", "Bishop", "Knight"],
value="Queen",
label="Promote pawns to",
)
movetime = gr.Slider(0.5, 8.0, value=3.0, step=0.5,
label="Think time per move (s)")
top_k = gr.Slider(3, 20, value=10, step=1,
label="Policy top-K explored per node")
max_depth = gr.Slider(1, 8, value=6, step=1, label="Max negamax depth")
gr.Markdown(
"The engine stops at whichever comes first: the think time, or the depth "
"limit. In endgames `search.py` extends the budget by 1.6× (author's rule)."
)
OUT = [board_img, history, selected, status_md, moves_md, analysis_out, legal_dd]
gr.Examples(
examples=[["e4"], ["d4"], ["Nf3"], ["c4"], ["g3"]],
inputs=[move_in],
outputs=OUT[:6] + [side_radio],
fn=play_opening,
cache_examples=True,
cache_mode="lazy",
label="Open with a classic first move and see how ChessMamba answers",
)
with gr.Accordion("A full ChessMamba game (author's replay, ends in checkmate)",
open=False):
gr.Video(value="replay.mp4", label="replay.mp4 — from the 4-0 match",
interactive=False)
gr.Markdown(
"Sources: model, weights, engine code and replay from "
f"[`{MODEL_ID}`](https://huggingface.co/{MODEL_ID}) (MIT). "
"`model.py`, `search.py` and `chess_io.py` are vendored **verbatim** from "
"that repo, so the moves you see here are the reference engine's."
)
board_img.select(
board_click,
inputs=[history, selected, side_radio, movetime, top_k, max_depth, promo_radio],
outputs=OUT,
)
play_btn.click(
play_move,
inputs=[move_in, legal_dd, history, side_radio, movetime, top_k, max_depth,
promo_radio],
outputs=OUT,
api_name="play_move",
).then(lambda: "", outputs=move_in)
move_in.submit(
play_move,
inputs=[move_in, legal_dd, history, side_radio, movetime, top_k, max_depth,
promo_radio],
outputs=OUT,
).then(lambda: "", outputs=move_in)
new_btn.click(
new_game,
inputs=[side_radio, movetime, top_k, max_depth],
outputs=OUT,
api_name="new_game",
)
side_radio.change(
new_game,
inputs=[side_radio, movetime, top_k, max_depth],
outputs=OUT,
)
undo_btn.click(undo, inputs=[history, side_radio, top_k], outputs=OUT)
if __name__ == "__main__":
demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS, show_error=True)