Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 14 additions & 14 deletions azlite_portfolio_clean.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,14 @@
import torch.nn as nn
import torch.optim as optim


# ----------------------------- Constant ---------------------------------------
PROMOTION_MAP = {
chess.QUEEN: 1,
chess.ROOK: 2,
chess.BISHOP: 3,
chess.KNIGHT: 4,
}
# ----------------------------- Config ---------------------------------------

DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
Expand Down Expand Up @@ -241,9 +249,7 @@ def run(
if mv.promotion is None:
promo_idxs.append(0)
else:
promo_map = {chess.QUEEN: 1, chess.ROOK: 2,
chess.BISHOP: 3, chess.KNIGHT: 4}
promo_idxs.append(promo_map.get(mv.promotion, 0))
promo_idxs.append(PROMOTION_MAP.get(mv.promotion, 0))
from_idx_t = torch.tensor(
from_idxs, dtype=torch.long, device=DEVICE)
to_idx_t = torch.tensor(to_idxs, dtype=torch.long, device=DEVICE)
Expand Down Expand Up @@ -338,9 +344,8 @@ def _simulate(self, board: chess.Board, root: MCTSNode) -> float:
if mv.promotion is None:
promo_idxs.append(0)
else:
promo_map = {chess.QUEEN: 1, chess.ROOK: 2,
chess.BISHOP: 3, chess.KNIGHT: 4}
promo_idxs.append(promo_map.get(mv.promotion, 0))
PROMOTION_MAP.get(mv.promotion, 0)

Copilot AI Sep 16, 2025

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Line 348 calls PROMOTION_MAP.get(mv.promotion, 0) but doesn't use the result. This appears to be a copy-paste error - line 348 should be removed as line 349 correctly appends the value to the list.

Suggested change
PROMOTION_MAP.get(mv.promotion, 0)

Copilot uses AI. Check for mistakes.
promo_idxs.append(PROMOTION_MAP.get(mv.promotion, 0))
logits = self.net.score_moves(
state_embed.squeeze(0),
torch.tensor(from_idxs, dtype=torch.long,
Expand Down Expand Up @@ -434,16 +439,11 @@ def self_play_episode(mcts: MCTS, max_moves: int = SELFPLAY_MAX_MOVES,
arr_state = board_to_tensor(board)
from_idxs = [mv.from_square for mv in legal]
to_idxs = [mv.to_square for mv in legal]
promo_map = {
chess.QUEEN: 1,
chess.ROOK: 2,
chess.BISHOP: 3,
chess.KNIGHT: 4,
}

promo_idxs = [
0 if mv.promotion is None else promo_map.get(mv.promotion, 0)
0 if mv.promotion is None else PROMOTION_MAP.get(mv.promotion, 0)
for mv in legal
]
]
examples.append(
SelfPlayExample(
state=arr_state,
Expand Down
61 changes: 61 additions & 0 deletions tests/test_promotion_mapping.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
import sys
import os

# Add parent directory to Python path so azlite_portfolio_clean can be imported
sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))

import chess
import numpy as np
from azlite_portfolio_clean import PROMOTION_MAP, self_play_episode

def test_promotion_map_constant():
"""Test that PROMOTION_MAP constant exists and has correct values"""
# Test the constant exists and has expected structure
assert PROMOTION_MAP is not None
assert isinstance(PROMOTION_MAP, dict)
assert len(PROMOTION_MAP) == 4

# Test specific mappings
assert PROMOTION_MAP[chess.QUEEN] == 1
assert PROMOTION_MAP[chess.ROOK] == 2
assert PROMOTION_MAP[chess.BISHOP] == 3
assert PROMOTION_MAP[chess.KNIGHT] == 4

def test_promotion_map_values():
"""Test that PROMOTION_MAP has correct chess piece mappings"""
expected = {
chess.QUEEN: 1,
chess.ROOK: 2,
chess.BISHOP: 3,
chess.KNIGHT: 4,
}
assert PROMOTION_MAP == expected

def test_promotion_map_usage():
"""Test that the promotion mapping logic works correctly"""

assert PROMOTION_MAP.get(chess.QUEEN, 0) == 1
assert PROMOTION_MAP.get(chess.ROOK, 0) == 2
assert PROMOTION_MAP.get(chess.BISHOP, 0) == 3
assert PROMOTION_MAP.get(chess.KNIGHT, 0) == 4
assert PROMOTION_MAP.get(None, 0) == 0

def test_promotion_indices_logic():
"""Test the promotion index calculation logic from self_play_episode"""

mock_moves = [
type('MockMove', (), {'promotion': chess.QUEEN})(),
type('MockMove', (), {'promotion': chess.ROOK})(),
type('MockMove', (), {'promotion': chess.BISHOP})(),
type('MockMove', (), {'promotion': chess.KNIGHT})(),
type('MockMove', (), {'promotion': None})(),
]


promo_idxs = [
0 if mv.promotion is None else PROMOTION_MAP.get(mv.promotion, 0)
for mv in mock_moves
]

expected = [1, 2, 3, 4, 0] # Queen, Rook, Bishop, Knight, None
assert promo_idxs == expected
Loading