| """Standalone ONNX inference for mini-chessformer-v1. |
| |
| This file is **self-contained**: it depends only on `chess`, `numpy`, and |
| `onnxruntime`. It carries its own copy of the board encoding, the 4272-move |
| table, and the legal-move mask, so HF consumers do not need to clone the |
| chessdb repository or import `engine.interfaces` / `experiments.chessformer_lite.encode`. |
| |
| Usage: |
| python inference.py --model mini-chessformer-v1.onnx --fens <list> |
| or: |
| from inference import ChessformerLiteONNX |
| eng = ChessformerLiteONNX("mini-chessformer-v1.onnx") |
| policy_logits, wdl_logits, best_move = eng.evaluate(board, contempt=0.0) |
| |
| Inputs (ONNX): |
| square_ids : int64 [B, 64] |
| state_features: float32 [B, 8] |
| contempt : float32 [B] |
| Outputs: |
| policy : float32 [B, 4272] raw logits (apply legal_mask + softmax) |
| wdl : float32 [B, 3] raw logits (win/draw/loss, mover POV) |
| """ |
|
|
| from __future__ import annotations |
|
|
| import argparse |
|
|
| import chess |
| import numpy as np |
| import onnxruntime as ort |
|
|
| |
| |
| |
| STATE_DIM = 8 |
| MOVE_SPACE = 4272 |
| N_BASE = 64 * 64 |
| N_PROMO = MOVE_SPACE - N_BASE |
|
|
| |
| _PIECE_OFFSET = { |
| chess.PAWN: 1, |
| chess.KNIGHT: 2, |
| chess.BISHOP: 3, |
| chess.ROOK: 4, |
| chess.QUEEN: 5, |
| chess.KING: 6, |
| } |
|
|
| _PROMO_PIECES = (chess.KNIGHT, chess.BISHOP, chess.ROOK, chess.QUEEN) |
|
|
|
|
| |
| |
| |
| |
| |
| |
| def _build_move_index() -> tuple[list, dict]: |
| moves: list[tuple[int, int, int | None]] = [] |
| for frm in range(64): |
| for to in range(64): |
| moves.append((frm, to, None)) |
| for frm in range(64): |
| fr, ff = chess.square_rank(frm), chess.square_file(frm) |
| for to in range(64): |
| tr, tf = chess.square_rank(to), chess.square_file(to) |
| white_promo = fr == 6 and tr == 7 |
| black_promo = fr == 1 and tr == 0 |
| if (white_promo or black_promo) and abs(ff - tf) <= 1: |
| for p in _PROMO_PIECES: |
| moves.append((frm, to, p)) |
| return moves, {m: i for i, m in enumerate(moves)} |
|
|
|
|
| _MOVES_LIST, _MOVE_TO_IDX = _build_move_index() |
| assert len(_MOVES_LIST) == MOVE_SPACE, len(_MOVES_LIST) |
|
|
|
|
| def move_to_index(move: chess.Move) -> int: |
| return _MOVE_TO_IDX[(move.from_square, move.to_square, move.promotion)] |
|
|
|
|
| def index_to_move(index: int) -> chess.Move: |
| frm, to, promo = _MOVES_LIST[index] |
| return chess.Move(frm, to, promotion=promo) |
|
|
|
|
| def legal_mask(board: chess.Board) -> np.ndarray: |
| """Return bool array shape (MOVE_SPACE,) True at legal move indices.""" |
| mask = np.zeros(MOVE_SPACE, dtype=bool) |
| for m in board.legal_moves: |
| mask[_MOVE_TO_IDX[(m.from_square, m.to_square, m.promotion)]] = True |
| assert mask.any(), "no legal moves (terminal position)" |
| return mask |
|
|
|
|
| |
| |
| |
| def board_to_square_ids(board: chess.Board) -> np.ndarray: |
| """int64 [64] piece ids. a1..h8, empty=0, white P..K=1..6, black P..K=7..12.""" |
| ids = np.zeros(64, dtype=np.int64) |
| for sq in chess.SQUARES: |
| piece = board.piece_at(sq) |
| if piece is None: |
| ids[sq] = 0 |
| else: |
| offset = _PIECE_OFFSET[piece.piece_type] |
| ids[sq] = offset if piece.color == chess.WHITE else offset + 6 |
| return ids |
|
|
|
|
| def board_to_state_features(board: chess.Board, repetition_count: int | None = None) -> np.ndarray: |
| """float32 [STATE_DIM] — see module docstring of the source package. |
| |
| Layout: |
| [0] side to move: 1.0 white / 0.0 black |
| [1:5] castling WK, WQ, BK, BQ as 0/1 |
| [5] ep file / 7, or -1 if none |
| [6] halfmove bucket min(clock//5, 9)/9 |
| [7] repetition 0/0.5/1.0 (from board.is_repetition or override) |
| """ |
| feat = np.empty(STATE_DIM, dtype=np.float32) |
| feat[0] = 1.0 if board.turn == chess.WHITE else 0.0 |
| feat[1] = 1.0 if board.has_kingside_castling_rights(chess.WHITE) else 0.0 |
| feat[2] = 1.0 if board.has_queenside_castling_rights(chess.WHITE) else 0.0 |
| feat[3] = 1.0 if board.has_kingside_castling_rights(chess.BLACK) else 0.0 |
| feat[4] = 1.0 if board.has_queenside_castling_rights(chess.BLACK) else 0.0 |
| if board.ep_square is not None: |
| feat[5] = chess.square_file(board.ep_square) / 7.0 |
| else: |
| feat[5] = -1.0 |
| feat[6] = min(board.halfmove_clock // 5, 9) / 9.0 |
| if repetition_count is not None: |
| prior = max(0, int(repetition_count) - 1) |
| rep = min(prior, 2) |
| elif board.is_repetition(3): |
| rep = 2 |
| elif board.is_repetition(2): |
| rep = 1 |
| else: |
| rep = 0 |
| feat[7] = rep / 2.0 |
| return feat |
|
|
|
|
| |
| |
| |
| class ChessformerLiteONNX: |
| """ONNX runtime wrapper for mini-chessformer-v1. Batched or single-board.""" |
|
|
| def __init__(self, model_path: str, providers: list[str] | None = None): |
| self.session = ort.InferenceSession( |
| model_path, |
| providers=providers or ["CPUExecutionProvider"], |
| ) |
| self.input_names = [i.name for i in self.session.get_inputs()] |
| self.output_names = [o.name for o in self.session.get_outputs()] |
| |
| assert self.input_names == ["square_ids", "state_features", "contempt"], self.input_names |
| assert self.output_names == ["policy", "wdl"], self.output_names |
|
|
| def evaluate_batch( |
| self, boards: list[chess.Board], contempt: float = 0.0 |
| ) -> tuple[np.ndarray, np.ndarray]: |
| """Return (policy_logits [N, 4272], wdl_logits [N, 3]) raw logits.""" |
| n = len(boards) |
| assert n >= 1 |
| sq = np.stack([board_to_square_ids(b) for b in boards]).astype(np.int64) |
| st = np.stack([board_to_state_features(b) for b in boards]).astype(np.float32) |
| c = np.full((n,), float(contempt), dtype=np.float32) |
| policy, wdl = self.session.run( |
| self.output_names, |
| {"square_ids": sq, "state_features": st, "contempt": c}, |
| ) |
| return policy, wdl |
|
|
| def evaluate( |
| self, board: chess.Board, contempt: float = 0.0 |
| ) -> tuple[np.ndarray, np.ndarray, chess.Move]: |
| """Return (policy_logits [4272], wdl_logits [3], best_legal_move).""" |
| pol_flat, wdl_flat = self.evaluate_batch([board], contempt=contempt) |
| logits = pol_flat[0] |
| mask = legal_mask(board) |
| masked = logits - 1e9 * (1.0 - mask) |
| best_idx = int(np.argmax(masked)) |
| return logits, wdl_flat[0], index_to_move(best_idx) |
|
|
|
|
| |
| |
| |
| def main() -> None: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--model", required=True, help="Path to mini-chessformer-v1.onnx") |
| parser.add_argument( |
| "--fens", |
| nargs="+", |
| default=[ |
| chess.STARTING_FEN, |
| "r1bqkbnr/pppp1ppp/2n5/4p3/4P3/5N2/PPPP1PPP/RNBQKB1R b KQkq - 2 2", |
| "8/P7/8/8/8/8/8/k6K w - - 0 1", |
| ], |
| help="FENs to evaluate (defaults: startpos, midgame, promotion)", |
| ) |
| parser.add_argument("--contempt", type=float, default=0.0) |
| args = parser.parse_args() |
|
|
| eng = ChessformerLiteONNX(args.model) |
| boards = [chess.Board(fen) for fen in args.fens] |
|
|
| |
| print(f"batched call: {len(boards)} boards, contempt={args.contempt}") |
| pol, wdl = eng.evaluate_batch(boards, contempt=args.contempt) |
| print(f" policy shape={pol.shape} dtype={pol.dtype}") |
| print(f" wdl shape={wdl.shape} dtype={wdl.dtype}") |
| assert pol.shape == (len(boards), MOVE_SPACE) |
| assert wdl.shape == (len(boards), 3) |
| assert np.all(np.isfinite(pol)) and np.all(np.isfinite(wdl)), "non-finite output" |
|
|
| |
| print() |
| for b in boards: |
| _, w_i, best = eng.evaluate(b, contempt=args.contempt) |
| legal = list(b.legal_moves) |
| assert best in legal, f"best {best} not in legal {legal}" |
| |
| probs = np.exp(w_i - w_i.max()) |
| probs /= probs.sum() |
| w, d, l = probs |
| print( |
| f" fen={b.fen() !r}\n" |
| f" best={b.san(best)} WDL(W/D/L)={w:.3f}/{d:.3f}/{l:.3f}" |
| ) |
|
|
| |
| |
| print() |
| print("single-board batch=1 check (boards[0]):") |
| pol1, wdl1, best1 = eng.evaluate(boards[0], contempt=args.contempt) |
| print(f" best={boards[0].san(best1)} policy[0] matches batch: {np.allclose(pol1, pol[0])}") |
| assert np.allclose(pol1, pol[0]), "batch=1 != first row of batch=N" |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|