Files
snake-python/snakes/bitboard_duel_search.py
T
daniel156161 cb6c8d4dc8 feat(snake): add adaptive adversarial search
- Share duel search contexts and transpositions across candidate moves.
- Add aspiration windows, principal variation ordering, and body caches.
- Model simultaneous multiplayer responses with a compact beam rollout.
- Adapt search depth and response breadth to the remaining deadline.
- Add a deterministic arena benchmark with optional JSON reporting.
- Expose search metrics, document benchmarking, and bump Prism to 1.2.0.
2026-08-01 19:26:19 +02:00

404 lines
14 KiB
Python

"""Deadline-aware simultaneous duel search using compact tuple bodies and bitboards."""
from __future__ import annotations
from collections.abc import Iterable
from dataclasses import dataclass
from time import perf_counter
from snakes.bitboard import BitBoard
Body = tuple[int, ...]
@dataclass(frozen=True, slots=True)
class DuelState:
my_body: Body
enemy_body: Body
food_bits: int
my_health: int
enemy_health: int
previous_hazard_bits: int
class BitboardDuelSearch:
"""Iterative-deepening paranoid minimax for a two-snake game.
The public API still accepts Battlesnake body dictionaries. Search nodes use
flat cell indices, immutable tuples, and integer masks to avoid allocation of
coordinate dictionaries and sets in the hot path.
"""
WIN = 100_000.0
LOSS = -100_000.0
def __init__(
self,
board: BitBoard,
food: Iterable[tuple[int, int]],
hazards: Iterable[tuple[int, int]],
hazard_count: dict[tuple[int, int], int],
hazard_damage: int,
deadline: float | None,
) -> None:
self.board = board
self.deadline = deadline
self.hazard_damage = hazard_damage
self.food_bits = board.set_to_bits(set(food))
self.hazard_bits = board.set_to_bits(set(hazards))
self.hazard_stacks = {
board.idx(x, y): count for (x, y), count in hazard_count.items()
}
self.transposition: dict[tuple[DuelState, int], tuple[float, str, int | None]] = {}
self.killer_moves: dict[int, int] = {}
self.history: dict[int, int] = {}
self._body_bits_cache: dict[Body, int] = {}
self.nodes = 0
self.cache_hits = 0
def body_from_dicts(self, body: list[dict]) -> Body:
return tuple(self.board.idx(seg["x"], seg["y"]) for seg in body)
def search(
self,
my_body: list[dict],
enemy_body: list[dict],
my_health: int,
enemy_health: int,
max_depth: int,
previous_hazards: Iterable[tuple[int, int]],
) -> tuple[float, int]:
state = DuelState(
my_body=self.body_from_dicts(my_body),
enemy_body=self.body_from_dicts(enemy_body),
food_bits=self.food_bits,
my_health=my_health,
enemy_health=enemy_health,
previous_hazard_bits=self.board.set_to_bits(set(previous_hazards)),
)
result = self._evaluate(state)
completed_depth = 0
for depth in range(1, max_depth + 1):
if self._out_of_time(5.0):
break
window = 80.0 if completed_depth else float("inf")
alpha, beta = result - window, result + window
value, completed = self._search(state, depth, alpha, beta)
if completed and window != float("inf") and (value <= alpha or value >= beta):
value, completed = self._search(state, depth, -float("inf"), float("inf"))
if not completed:
break
result = value
completed_depth = depth
return result, completed_depth
def search_candidate(
self,
my_body: list[dict],
enemy_body: list[dict],
my_target: tuple[int, int],
my_health: int,
enemy_health: int,
max_depth: int,
previous_hazards: Iterable[tuple[int, int]],
) -> tuple[float, int]:
"""Evaluate one selected move against every simultaneous enemy reply.
``max_depth`` counts the selected root turn, so a completed depth of one
means all opponent replies to that move were resolved.
"""
state = DuelState(
my_body=self.body_from_dicts(my_body),
enemy_body=self.body_from_dicts(enemy_body),
food_bits=self.food_bits,
my_health=my_health,
enemy_health=enemy_health,
previous_hazard_bits=self.board.set_to_bits(set(previous_hazards)),
)
target_idx = self.board.idx(my_target[0], my_target[1])
if not self.board.neighbors_of(state.my_body[0]) & (1 << target_idx):
return self.LOSS, 0
result = self._evaluate(state)
completed_depth = 0
for depth in range(1, max_depth + 1):
if self._out_of_time(5.0):
break
window = 80.0 if completed_depth else float("inf")
alpha, beta = result - window, result + window
value, completed = self._search_selected_move(state, target_idx, depth, alpha, beta)
if completed and window != float("inf") and (value <= alpha or value >= beta):
value, completed = self._search_selected_move(
state, target_idx, depth, -float("inf"), float("inf")
)
if not completed:
break
result = value
completed_depth = depth
return result, completed_depth
def search_depth(
self,
my_body: list[dict],
enemy_body: list[dict],
my_health: int,
enemy_health: int,
depth: int,
previous_hazards: Iterable[tuple[int, int]],
) -> float:
state = DuelState(
my_body=self.body_from_dicts(my_body),
enemy_body=self.body_from_dicts(enemy_body),
food_bits=self.food_bits,
my_health=my_health,
enemy_health=enemy_health,
previous_hazard_bits=self.board.set_to_bits(set(previous_hazards)),
)
value, _ = self._search(state, depth, -float("inf"), float("inf"))
return value
def _search_selected_move(
self,
state: DuelState,
my_target: int,
depth: int,
alpha: float,
beta: float,
) -> tuple[float, bool]:
"""Resolve the selected root move with the opponent on the same turn."""
self.nodes += 1
if self._out_of_time():
return self._evaluate(state), False
enemy_moves = self._candidate_targets(state.enemy_body)
if not enemy_moves:
return self.WIN + depth, True
enemy_moves = self._ordered_moves(enemy_moves, state, depth, False)
worst = float("inf")
for enemy_target in enemy_moves:
if self._out_of_time():
return (worst if worst != float("inf") else self._evaluate(state)), False
child, terminal = self._advance(state, my_target, enemy_target)
if terminal is not None:
value = terminal
completed = True
elif depth <= 1:
value = self._evaluate(child)
completed = True
else:
value, completed = self._search(child, depth - 1, alpha, beta)
if not completed:
return (worst if worst != float("inf") else value), False
worst = min(worst, value)
beta = min(beta, worst)
if beta <= alpha:
break
return worst, True
def _search(self, state: DuelState, depth: int, alpha: float, beta: float) -> tuple[float, bool]:
self.nodes += 1
if self._out_of_time():
return self._evaluate(state), False
if depth <= 0:
return self._evaluate(state), True
cache_key = (state, depth)
original_alpha, original_beta = alpha, beta
cached = self.transposition.get(cache_key)
if cached is not None:
self.cache_hits += 1
cached_value, bound, preferred_move = cached
if bound == "exact":
return cached_value, True
if bound == "lower":
alpha = max(alpha, cached_value)
else:
beta = min(beta, cached_value)
if alpha >= beta:
return cached_value, True
my_moves = self._candidate_targets(state.my_body)
enemy_moves = self._candidate_targets(state.enemy_body)
if not my_moves:
return self.LOSS - depth, True
if not enemy_moves:
return self.WIN + depth, True
my_moves = self._ordered_moves(my_moves, state, depth, True, preferred_move if cached is not None else None)
enemy_moves = self._ordered_moves(enemy_moves, state, depth, False)
best = -float("inf")
best_move: int | None = None
for my_target in my_moves:
worst = float("inf")
for enemy_target in enemy_moves:
if self._out_of_time():
return (best if best != -float("inf") else self._evaluate(state)), False
child, terminal = self._advance(state, my_target, enemy_target)
if terminal is not None:
value = terminal
completed = True
else:
value, completed = self._search(child, depth - 1, alpha, beta)
if not completed:
return (best if best != -float("inf") else value), False
worst = min(worst, value)
if worst <= alpha:
self.killer_moves[depth] = my_target
self.history[my_target] = self.history.get(my_target, 0) + depth * depth
break
if worst > best:
best = worst
best_move = my_target
alpha = max(alpha, best)
if alpha >= beta:
break
if best <= original_alpha:
bound = "upper"
elif best >= original_beta:
bound = "lower"
else:
bound = "exact"
self.transposition[cache_key] = (best, bound, best_move)
return best, True
def _advance(self, state: DuelState, my_target: int, enemy_target: int) -> tuple[DuelState, float | None]:
my_ate = bool((1 << my_target) & state.food_bits)
enemy_ate = bool((1 << enemy_target) & state.food_bits)
my_body = self._advance_body(state.my_body, my_target, my_ate)
enemy_body = self._advance_body(state.enemy_body, enemy_target, enemy_ate)
my_dead = my_target in my_body[1:] or my_target in enemy_body[1:]
enemy_dead = enemy_target in enemy_body[1:] or enemy_target in my_body[1:]
if my_target == enemy_target:
if len(my_body) <= len(enemy_body):
my_dead = True
if len(enemy_body) <= len(my_body):
enemy_dead = True
my_health = 100 if my_ate else state.my_health - 1
enemy_health = 100 if enemy_ate else state.enemy_health - 1
if not my_ate:
my_health -= self._hazard_cost(my_target, state.previous_hazard_bits)
if not enemy_ate:
enemy_health -= self._hazard_cost(enemy_target, state.previous_hazard_bits)
my_dead = my_dead or my_health <= 0
enemy_dead = enemy_dead or enemy_health <= 0
if my_dead and enemy_dead:
return state, -500.0
if my_dead:
return state, self.LOSS
if enemy_dead:
return state, self.WIN
eaten_bits = 0
if my_ate:
eaten_bits |= 1 << my_target
if enemy_ate:
eaten_bits |= 1 << enemy_target
child = DuelState(
my_body=my_body,
enemy_body=enemy_body,
food_bits=state.food_bits & ~eaten_bits,
my_health=my_health,
enemy_health=enemy_health,
previous_hazard_bits=self.hazard_bits,
)
return child, None
def _candidate_targets(self, body: Body) -> list[int]:
"""Return in-bounds targets; `_advance` resolves simultaneous collisions.
Delaying occupancy checks until both targets and food growth are known is
essential: whether either tail vacates depends on that snake eating.
"""
return list(self._iter_bits(self.board.neighbors_of(body[0])))
def _ordered_moves(
self, moves: list[int], state: DuelState, depth: int, mine: bool,
preferred: int | None = None,
) -> list[int]:
body = state.my_body if mine else state.enemy_body
other = state.enemy_body if mine else state.my_body
killer = self.killer_moves.get(depth)
center_x = (self.board.width - 1) / 2.0
center_y = (self.board.height - 1) / 2.0
def score(target: int) -> tuple[float, int]:
x, y = self.board.coord(target)
food_bonus = 200.0 if (1 << target) & state.food_bits else 0.0
space = self.board.flood_count(target, (self._body_bits(body[1:]) | self._body_bits(other[1:])) & ~(1 << target))
center = -(abs(x - center_x) + abs(y - center_y))
preferred_bonus = 20_000.0 if target == preferred else 0.0
killer_bonus = 10_000.0 if target == killer else 0.0
return preferred_bonus + killer_bonus + self.history.get(target, 0) + food_bonus + space * 2.0 + center, -target
# Our strongest-looking moves first; enemy ordering uses the same quality
# estimate because dangerous enemy replies tend to gain space and food.
return sorted(moves, key=score, reverse=True)
def _evaluate(self, state: DuelState) -> float:
my_blocked = self._body_bits(state.my_body[1:]) | self._body_bits(state.enemy_body[1:])
my_head, enemy_head = state.my_body[0], state.enemy_body[0]
my_space = self.board.flood_count(my_head, my_blocked)
enemy_space = self.board.flood_count(enemy_head, my_blocked)
my_liberties = self.board.open_neighbor_count(my_head, my_blocked)
enemy_liberties = self.board.open_neighbor_count(enemy_head, my_blocked)
territory = self.board.territory(my_head, [enemy_head], my_blocked)
my_tail_path = self.board.path_distance(my_head, state.my_body[-1], my_blocked)
enemy_tail_path = self.board.path_distance(enemy_head, state.enemy_body[-1], my_blocked)
tail_score = (12.0 if my_tail_path is not None else -24.0) - (12.0 if enemy_tail_path is not None else -24.0)
my_hazard = self._hazard_cost(my_head, state.previous_hazard_bits)
enemy_hazard = self._hazard_cost(enemy_head, state.previous_hazard_bits)
length_score = (len(state.my_body) - len(state.enemy_body)) * 20.0
health_score = (state.my_health - state.enemy_health) * 0.18
forced_score = (my_liberties > 1) * 10.0 - (enemy_liberties > 1) * 10.0
return (
(my_space - enemy_space) * 1.5 + territory * 1.2
+ (my_liberties - enemy_liberties) * 14.0 + length_score + health_score
+ tail_score + forced_score + (enemy_hazard - my_hazard) * 0.8
)
def _hazard_cost(self, target: int, previous_hazard_bits: int) -> int:
bit = 1 << target
if not (bit & self.hazard_bits & previous_hazard_bits):
return 0
return self.hazard_damage * self.hazard_stacks.get(target, 1)
@staticmethod
def _advance_body(body: Body, target: int, ate: bool) -> Body:
return (target,) + body if ate else (target,) + body[:-1]
@staticmethod
def _tail_stacked(body: Body) -> bool:
return len(body) >= 2 and body[-1] == body[-2]
def _body_bits(self, body: Body) -> int:
cached = self._body_bits_cache.get(body)
if cached is not None:
return cached
bits = 0
for cell in body:
bits |= 1 << cell
if len(self._body_bits_cache) < 16_384:
self._body_bits_cache[body] = bits
return bits
@staticmethod
def _iter_bits(bits: int):
while bits:
bit = bits & -bits
yield bit.bit_length() - 1
bits ^= bit
def _out_of_time(self, reserve_ms: float = 0.0) -> bool:
if self.deadline is None:
return False
return perf_counter() + reserve_ms / 1000.0 >= self.deadline