From 3b2ffa28d4e1ea49b320756b817b4c70a8cd4334 Mon Sep 17 00:00:00 2001 From: "mattlamb227@gmail.com" Date: Wed, 12 Aug 2026 23:39:52 -0400 Subject: [PATCH] Optimize: Cache player occupied squares for faster corner validation --- src/blokus_gym/core/board.py | 10 +++++++ src/blokus_gym/core/game.py | 53 +++++++++++++++++------------------- 2 files changed, 35 insertions(+), 28 deletions(-) diff --git a/src/blokus_gym/core/board.py b/src/blokus_gym/core/board.py index 5dccba7..3760045 100644 --- a/src/blokus_gym/core/board.py +++ b/src/blokus_gym/core/board.py @@ -16,11 +16,15 @@ class Board: def __init__(self, size: int): self.size = size self.grid = np.zeros((size, size), dtype=np.int8) + self._occupied_cache: dict[int, set[tuple[int, int]]] = {} # player_idx -> occupied squares def place(self, player_idx: int, squares: list[tuple[int, int]]) -> None: """Place a player's piece on the board.""" for x, y in squares: self.grid[y, x] = player_idx + # Update occupied cache incrementally instead of clearing + if player_idx in self._occupied_cache: + self._occupied_cache[player_idx].update(squares) def clear(self) -> None: """Reset the board to empty.""" @@ -50,6 +54,12 @@ class Board: ys, xs = np.where(self.grid == player_idx) return [(int(x), int(y)) for x, y in zip(xs, ys, strict=True)] + def get_player_occupied(self, player_idx: int) -> set[tuple[int, int]]: + """Get all squares occupied by a player as a set for fast lookup.""" + if player_idx not in self._occupied_cache: + self._occupied_cache[player_idx] = set(self.get_player_squares(player_idx)) + return self._occupied_cache[player_idx] + def get_player_corners(self, player_idx: int) -> set[tuple[int, int]]: """Get all corner cells adjacent to a player's pieces. diff --git a/src/blokus_gym/core/game.py b/src/blokus_gym/core/game.py index cdfa523..644bdb8 100644 --- a/src/blokus_gym/core/game.py +++ b/src/blokus_gym/core/game.py @@ -206,12 +206,12 @@ class BlokusGame: # Rule 4: Corner rule (must touch same-color at a corner) if player.has_started: placed_corners = self._get_placed_corners(move) - touches_corner = False - for cx, cy in placed_corners: - if self.board.in_bounds(cx, cy): - if self.board.get_cell(cx, cy) == player_board_idx: - touches_corner = True - break + # Use cached occupied squares from board for fast lookup + player_occupied = self.board.get_player_occupied(player_idx) + touches_corner = any( + self.board.in_bounds(cx, cy) and (cx, cy) in player_occupied + for cx, cy in placed_corners + ) if not touches_corner: return False @@ -238,49 +238,46 @@ class BlokusGame: return mask player_board_idx = player_idx + self.PLAYER_OFFSET + + # Pre-compute player's occupied squares ONCE (cached) + player_occupied = self.board.get_player_occupied(player_idx) for piece_id in available_piece_ids: for action_idx in self._actions_by_piece[piece_id]: move = self._action_moves[action_idx] placed_corners = self._get_placed_corners(move) - touches_corner = False - for cx, cy in placed_corners: - if self.board.in_bounds(cx, cy): - if self.board.get_cell(cx, cy) == player_board_idx: - touches_corner = True - break + # Early exit: check corner touch using cached occupied squares + touches_corner = any( + self.board.in_bounds(cx, cy) and (cx, cy) in player_occupied + for cx, cy in placed_corners + ) if not touches_corner: - continue + continue # Skip expensive overlap/edge checks placed_squares = self._get_placed_squares(move) - valid = True - for x, y in placed_squares: - if not self.board.in_bounds(x, y): - valid = False - break - if not valid: + # Check bounds + if not all(self.board.in_bounds(x, y) for x, y in placed_squares): continue + # Check overlap if self.board.has_overlap(placed_squares): continue + # Check edge adjacency (no same-color edges) + edge_invalid = False for x, y in placed_squares: - edge_invalid = False for dx, dy in [(-1, 0), (1, 0), (0, -1), (0, 1)]: nx, ny = x + dx, y + dy - if self.board.in_bounds(nx, ny): - if self.board.get_cell(nx, ny) == player_board_idx: - edge_invalid = True - break + if self.board.in_bounds(nx, ny) and self.board.get_cell(nx, ny) == player_board_idx: + edge_invalid = True + break if edge_invalid: - valid = False break - if not valid: - continue - mask[action_idx] = True + if not edge_invalid: + mask[action_idx] = True return mask