Deep CFR 기초 모듈 추가
This commit is contained in:
@@ -10,6 +10,7 @@ wheels/
|
|||||||
# Cython-generated sources
|
# Cython-generated sources
|
||||||
src/coolrl_lost_cities/games/classic/game.c
|
src/coolrl_lost_cities/games/classic/game.c
|
||||||
src/coolrl_lost_cities/games/classic/engines/fast.c
|
src/coolrl_lost_cities/games/classic/engines/fast.c
|
||||||
|
src/coolrl_lost_cities/games/classic/deep_cfr/*.c
|
||||||
|
|
||||||
# Rust build output
|
# Rust build output
|
||||||
target/
|
target/
|
||||||
|
|||||||
@@ -46,6 +46,10 @@ include = ["coolrl_lost_cities*"]
|
|||||||
"*.pxd",
|
"*.pxd",
|
||||||
"*.pyx",
|
"*.pyx",
|
||||||
]
|
]
|
||||||
|
"coolrl_lost_cities.games.classic.deep_cfr" = [
|
||||||
|
"*.pxd",
|
||||||
|
"*.pyx",
|
||||||
|
]
|
||||||
|
|
||||||
[tool.ruff]
|
[tool.ruff]
|
||||||
line-length = 100
|
line-length = 100
|
||||||
|
|||||||
@@ -14,6 +14,14 @@ extensions = cythonize(
|
|||||||
"coolrl_lost_cities.games.classic.game",
|
"coolrl_lost_cities.games.classic.game",
|
||||||
["src/coolrl_lost_cities/games/classic/game.pyx"],
|
["src/coolrl_lost_cities/games/classic/game.pyx"],
|
||||||
),
|
),
|
||||||
|
Extension(
|
||||||
|
"coolrl_lost_cities.games.classic.deep_cfr.cfr_math",
|
||||||
|
["src/coolrl_lost_cities/games/classic/deep_cfr/cfr_math.pyx"],
|
||||||
|
),
|
||||||
|
Extension(
|
||||||
|
"coolrl_lost_cities.games.classic.deep_cfr.encoding",
|
||||||
|
["src/coolrl_lost_cities/games/classic/deep_cfr/encoding.pyx"],
|
||||||
|
),
|
||||||
],
|
],
|
||||||
language_level=3,
|
language_level=3,
|
||||||
compiler_directives={
|
compiler_directives={
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
"""Minimal Deep CFR building blocks for Lost Cities classic."""
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
cdef int regret_matching_c(
|
||||||
|
const float* advantages,
|
||||||
|
const unsigned char* legal,
|
||||||
|
int n,
|
||||||
|
float epsilon,
|
||||||
|
float* out_policy,
|
||||||
|
) noexcept
|
||||||
|
|
||||||
|
cdef int normalize_legal_policy_c(
|
||||||
|
const float* weights,
|
||||||
|
const unsigned char* legal,
|
||||||
|
int n,
|
||||||
|
float* out_policy,
|
||||||
|
) noexcept
|
||||||
|
|
||||||
|
cdef int sample_policy_c(
|
||||||
|
const float* policy,
|
||||||
|
int n,
|
||||||
|
double random_value,
|
||||||
|
) noexcept
|
||||||
|
|
||||||
@@ -0,0 +1,146 @@
|
|||||||
|
# cython: language_level=3, boundscheck=False, wraparound=False, cdivision=True, initializedcheck=False
|
||||||
|
"""Small Cython math helpers for future Deep CFR traversal."""
|
||||||
|
|
||||||
|
|
||||||
|
cdef int regret_matching_c(
|
||||||
|
const float* advantages,
|
||||||
|
const unsigned char* legal,
|
||||||
|
int n,
|
||||||
|
float epsilon,
|
||||||
|
float* out_policy,
|
||||||
|
) noexcept:
|
||||||
|
cdef int i
|
||||||
|
cdef int legal_count = 0
|
||||||
|
cdef float positive
|
||||||
|
cdef float total = 0.0
|
||||||
|
if n <= 0:
|
||||||
|
return 0
|
||||||
|
for i in range(n):
|
||||||
|
if legal[i] != 0:
|
||||||
|
legal_count += 1
|
||||||
|
positive = advantages[i] if advantages[i] > 0.0 else 0.0
|
||||||
|
out_policy[i] = positive
|
||||||
|
total += positive
|
||||||
|
else:
|
||||||
|
out_policy[i] = 0.0
|
||||||
|
if legal_count <= 0:
|
||||||
|
return 0
|
||||||
|
if total > epsilon:
|
||||||
|
for i in range(n):
|
||||||
|
out_policy[i] = out_policy[i] / total
|
||||||
|
return legal_count
|
||||||
|
for i in range(n):
|
||||||
|
out_policy[i] = 1.0 / legal_count if legal[i] != 0 else 0.0
|
||||||
|
return legal_count
|
||||||
|
|
||||||
|
|
||||||
|
cdef int normalize_legal_policy_c(
|
||||||
|
const float* weights,
|
||||||
|
const unsigned char* legal,
|
||||||
|
int n,
|
||||||
|
float* out_policy,
|
||||||
|
) noexcept:
|
||||||
|
cdef int i
|
||||||
|
cdef int legal_count = 0
|
||||||
|
cdef float value
|
||||||
|
cdef float total = 0.0
|
||||||
|
if n <= 0:
|
||||||
|
return 0
|
||||||
|
for i in range(n):
|
||||||
|
if legal[i] != 0:
|
||||||
|
legal_count += 1
|
||||||
|
value = weights[i] if weights[i] > 0.0 else 0.0
|
||||||
|
out_policy[i] = value
|
||||||
|
total += value
|
||||||
|
else:
|
||||||
|
out_policy[i] = 0.0
|
||||||
|
if legal_count <= 0:
|
||||||
|
return 0
|
||||||
|
if total > 0.0:
|
||||||
|
for i in range(n):
|
||||||
|
out_policy[i] = out_policy[i] / total
|
||||||
|
return legal_count
|
||||||
|
for i in range(n):
|
||||||
|
out_policy[i] = 1.0 / legal_count if legal[i] != 0 else 0.0
|
||||||
|
return legal_count
|
||||||
|
|
||||||
|
|
||||||
|
cdef int sample_policy_c(
|
||||||
|
const float* policy,
|
||||||
|
int n,
|
||||||
|
double random_value,
|
||||||
|
) noexcept:
|
||||||
|
cdef int i
|
||||||
|
cdef int fallback = -1
|
||||||
|
cdef double cumulative = 0.0
|
||||||
|
cdef double r = random_value
|
||||||
|
if n <= 0:
|
||||||
|
return -1
|
||||||
|
if r < 0.0:
|
||||||
|
r = 0.0
|
||||||
|
elif r >= 1.0:
|
||||||
|
r = 0.9999999999999999
|
||||||
|
for i in range(n):
|
||||||
|
if policy[i] > 0.0:
|
||||||
|
fallback = i
|
||||||
|
cumulative += policy[i]
|
||||||
|
if r < cumulative:
|
||||||
|
return i
|
||||||
|
return fallback
|
||||||
|
|
||||||
|
|
||||||
|
def regret_matching(advantages, legal_mask, float epsilon=1.0e-8):
|
||||||
|
cdef float[::1] adv_view
|
||||||
|
cdef unsigned char[::1] legal_view
|
||||||
|
cdef float[::1] out_view
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
adv = np.ascontiguousarray(advantages, dtype=np.float32)
|
||||||
|
legal = np.ascontiguousarray(legal_mask, dtype=np.uint8)
|
||||||
|
if adv.ndim != 1 or legal.ndim != 1:
|
||||||
|
raise ValueError("advantages and legal_mask must be one-dimensional")
|
||||||
|
if adv.shape[0] != legal.shape[0]:
|
||||||
|
raise ValueError("advantages and legal_mask must have the same length")
|
||||||
|
out = np.empty_like(adv)
|
||||||
|
if adv.shape[0] == 0:
|
||||||
|
return out
|
||||||
|
adv_view = adv
|
||||||
|
legal_view = legal
|
||||||
|
out_view = out
|
||||||
|
regret_matching_c(&adv_view[0], &legal_view[0], adv.shape[0], epsilon, &out_view[0])
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_legal_policy(weights, legal_mask):
|
||||||
|
cdef float[::1] values_view
|
||||||
|
cdef unsigned char[::1] legal_view
|
||||||
|
cdef float[::1] out_view
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
values = np.ascontiguousarray(weights, dtype=np.float32)
|
||||||
|
legal = np.ascontiguousarray(legal_mask, dtype=np.uint8)
|
||||||
|
if values.ndim != 1 or legal.ndim != 1:
|
||||||
|
raise ValueError("weights and legal_mask must be one-dimensional")
|
||||||
|
if values.shape[0] != legal.shape[0]:
|
||||||
|
raise ValueError("weights and legal_mask must have the same length")
|
||||||
|
out = np.empty_like(values)
|
||||||
|
if values.shape[0] == 0:
|
||||||
|
return out
|
||||||
|
values_view = values
|
||||||
|
legal_view = legal
|
||||||
|
out_view = out
|
||||||
|
normalize_legal_policy_c(&values_view[0], &legal_view[0], values.shape[0], &out_view[0])
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def sample_policy(policy, double random_value):
|
||||||
|
cdef float[::1] values_view
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
values = np.ascontiguousarray(policy, dtype=np.float32)
|
||||||
|
if values.ndim != 1:
|
||||||
|
raise ValueError("policy must be one-dimensional")
|
||||||
|
if values.shape[0] == 0:
|
||||||
|
return -1
|
||||||
|
values_view = values
|
||||||
|
return sample_policy_c(&values_view[0], values.shape[0], random_value)
|
||||||
@@ -0,0 +1,6 @@
|
|||||||
|
from coolrl_lost_cities.games.classic.game cimport GameState
|
||||||
|
|
||||||
|
|
||||||
|
cdef int input_dim_c(GameState state) noexcept
|
||||||
|
cdef int encode_info_state_c(GameState state, int player, float* out) except -1
|
||||||
|
|
||||||
@@ -0,0 +1,72 @@
|
|||||||
|
# cython: language_level=3, boundscheck=False, wraparound=False, cdivision=True, initializedcheck=False
|
||||||
|
"""Minimal deterministic information-state encoding for Deep CFR scaffolding."""
|
||||||
|
|
||||||
|
from coolrl_lost_cities.games.classic.game cimport GameState
|
||||||
|
|
||||||
|
|
||||||
|
cdef int input_dim_c(GameState state) noexcept:
|
||||||
|
return 5 + state.hand_size * 3 + 2 * state.hand_size + 1 + state.n_colors
|
||||||
|
|
||||||
|
|
||||||
|
cdef int encode_info_state_c(GameState state, int player, float* out) except -1:
|
||||||
|
cdef int idx = 0
|
||||||
|
cdef int slot
|
||||||
|
cdef int card
|
||||||
|
cdef int color
|
||||||
|
cdef int rank
|
||||||
|
cdef int action_size = 2 * state.hand_size + 1 + state.n_colors
|
||||||
|
cdef int action_count
|
||||||
|
cdef int actions[64]
|
||||||
|
cdef int i
|
||||||
|
|
||||||
|
if player < 0 or player > 1:
|
||||||
|
raise ValueError(f"invalid player: {player}")
|
||||||
|
if action_size > 64:
|
||||||
|
raise ValueError("action_size exceeds fixed encoding action buffer")
|
||||||
|
|
||||||
|
out[idx] = 1.0 if state.phase_id == 0 else 0.0
|
||||||
|
idx += 1
|
||||||
|
out[idx] = 1.0 if state.phase_id == 1 else 0.0
|
||||||
|
idx += 1
|
||||||
|
out[idx] = <float>state.current_player
|
||||||
|
idx += 1
|
||||||
|
out[idx] = <float>player
|
||||||
|
idx += 1
|
||||||
|
out[idx] = <float>state.deck_len / <float>state.total_cards
|
||||||
|
idx += 1
|
||||||
|
|
||||||
|
for slot in range(state.hand_size):
|
||||||
|
if slot < state.hand_lens[player]:
|
||||||
|
card = state.hand_cards[state._hand_index(player, slot)]
|
||||||
|
color = state._card_color(card)
|
||||||
|
rank = state._card_rank(card)
|
||||||
|
out[idx] = 1.0
|
||||||
|
out[idx + 1] = <float>(color + 1) / <float>state.n_colors
|
||||||
|
out[idx + 2] = <float>rank / <float>state.n_ranks
|
||||||
|
else:
|
||||||
|
out[idx] = 0.0
|
||||||
|
out[idx + 1] = 0.0
|
||||||
|
out[idx + 2] = 0.0
|
||||||
|
idx += 3
|
||||||
|
|
||||||
|
for i in range(action_size):
|
||||||
|
out[idx + i] = 0.0
|
||||||
|
action_count = state._unified_legal_actions_c(actions)
|
||||||
|
for i in range(action_count):
|
||||||
|
out[idx + actions[i]] = 1.0
|
||||||
|
idx += action_size
|
||||||
|
return idx
|
||||||
|
|
||||||
|
|
||||||
|
def input_dim(GameState state) -> int:
|
||||||
|
return input_dim_c(state)
|
||||||
|
|
||||||
|
|
||||||
|
def encode_info_state(GameState state, int player):
|
||||||
|
cdef float[::1] out_view
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
out = np.empty(input_dim_c(state), dtype=np.float32)
|
||||||
|
out_view = out
|
||||||
|
encode_info_state_c(state, player, &out_view[0])
|
||||||
|
return out
|
||||||
@@ -73,6 +73,7 @@ cdef class GameState:
|
|||||||
cpdef int push_action(self, int action_id)
|
cpdef int push_action(self, int action_id)
|
||||||
cpdef int push_unified_action(self, int action_id)
|
cpdef int push_unified_action(self, int action_id)
|
||||||
cpdef int pop_action(self)
|
cpdef int pop_action(self)
|
||||||
|
cpdef swap_deck_cards(self, int left, int right)
|
||||||
cpdef bint can_play_encoded_card(self, int player, int card)
|
cpdef bint can_play_encoded_card(self, int player, int card)
|
||||||
cpdef int last_numeric_rank(self, int player, int color)
|
cpdef int last_numeric_rank(self, int player, int color)
|
||||||
cpdef int expedition_score(self, int player, int color)
|
cpdef int expedition_score(self, int player, int color)
|
||||||
@@ -89,6 +90,7 @@ cdef class GameState:
|
|||||||
cdef void _ensure_undo_capacity_c(self) except *
|
cdef void _ensure_undo_capacity_c(self) except *
|
||||||
cdef int _push_action_c(self, int action_id) except *
|
cdef int _push_action_c(self, int action_id) except *
|
||||||
cdef int _pop_action_c(self) except *
|
cdef int _pop_action_c(self) except *
|
||||||
|
cdef void _swap_deck_cards_c(self, int left, int right) except *
|
||||||
cdef object _undo_to_tuple(self, UndoRecord* undo)
|
cdef object _undo_to_tuple(self, UndoRecord* undo)
|
||||||
cdef void _tuple_to_undo(self, object data, UndoRecord* undo) except *
|
cdef void _tuple_to_undo(self, object data, UndoRecord* undo) except *
|
||||||
cdef void _apply_card_action(self, int action_id) except *
|
cdef void _apply_card_action(self, int action_id) except *
|
||||||
|
|||||||
@@ -737,6 +737,9 @@ cdef class GameState:
|
|||||||
raise ValueError("undo stack is empty")
|
raise ValueError("undo stack is empty")
|
||||||
return self._pop_action_c()
|
return self._pop_action_c()
|
||||||
|
|
||||||
|
cpdef swap_deck_cards(self, int left, int right):
|
||||||
|
self._swap_deck_cards_c(left, right)
|
||||||
|
|
||||||
cpdef bint can_play_encoded_card(self, int player, int card):
|
cpdef bint can_play_encoded_card(self, int player, int card):
|
||||||
cdef int color = self._card_color(card)
|
cdef int color = self._card_color(card)
|
||||||
cdef int rank = self._card_rank(card)
|
cdef int rank = self._card_rank(card)
|
||||||
@@ -1039,6 +1042,18 @@ cdef class GameState:
|
|||||||
self._undo_action_c(&self.undo_stack[self.undo_stack_len])
|
self._undo_action_c(&self.undo_stack[self.undo_stack_len])
|
||||||
return action_id
|
return action_id
|
||||||
|
|
||||||
|
cdef void _swap_deck_cards_c(self, int left, int right) except *:
|
||||||
|
cdef int tmp
|
||||||
|
if left < 0 or left >= self.deck_len:
|
||||||
|
raise IndexError(f"deck index out of range: {left}")
|
||||||
|
if right < 0 or right >= self.deck_len:
|
||||||
|
raise IndexError(f"deck index out of range: {right}")
|
||||||
|
if left == right:
|
||||||
|
return
|
||||||
|
tmp = self.deck_cards[left]
|
||||||
|
self.deck_cards[left] = self.deck_cards[right]
|
||||||
|
self.deck_cards[right] = tmp
|
||||||
|
|
||||||
cdef object _undo_to_tuple(self, UndoRecord* undo):
|
cdef object _undo_to_tuple(self, UndoRecord* undo):
|
||||||
return (
|
return (
|
||||||
"card" if undo.phase_id == _phase_card() else "draw",
|
"card" if undo.phase_id == _phase_card() else "draw",
|
||||||
|
|||||||
@@ -0,0 +1,95 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pytest
|
||||||
|
from coolrl_lost_cities.games.classic.deep_cfr.cfr_math import (
|
||||||
|
normalize_legal_policy,
|
||||||
|
regret_matching,
|
||||||
|
sample_policy,
|
||||||
|
)
|
||||||
|
from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state, input_dim
|
||||||
|
from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig
|
||||||
|
|
||||||
|
|
||||||
|
def test_swap_deck_cards_swaps_internal_deck_order_and_validates_bounds() -> None:
|
||||||
|
state = GameState.new_game(LostCitiesConfig(seed=7))
|
||||||
|
before = state.to_snapshot()
|
||||||
|
first = before["deck"][0]
|
||||||
|
last = before["deck"][-1]
|
||||||
|
|
||||||
|
state.swap_deck_cards(0, len(before["deck"]) - 1)
|
||||||
|
|
||||||
|
after = state.to_snapshot()
|
||||||
|
assert after["deck"][0] == last
|
||||||
|
assert after["deck"][-1] == first
|
||||||
|
state.validate_invariants()
|
||||||
|
|
||||||
|
state.swap_deck_cards(0, len(before["deck"]) - 1)
|
||||||
|
assert state.to_snapshot() == before
|
||||||
|
|
||||||
|
with pytest.raises(IndexError, match="deck index out of range"):
|
||||||
|
state.swap_deck_cards(-1, 0)
|
||||||
|
with pytest.raises(IndexError, match="deck index out of range"):
|
||||||
|
state.swap_deck_cards(0, len(before["deck"]))
|
||||||
|
|
||||||
|
|
||||||
|
def test_regret_matching_uses_positive_legal_regrets() -> None:
|
||||||
|
policy = regret_matching(
|
||||||
|
np.asarray([1.0, -2.0, 3.0, 5.0], dtype=np.float32),
|
||||||
|
np.asarray([True, True, False, True]),
|
||||||
|
)
|
||||||
|
|
||||||
|
np.testing.assert_allclose(policy, [1.0 / 6.0, 0.0, 0.0, 5.0 / 6.0])
|
||||||
|
|
||||||
|
|
||||||
|
def test_regret_matching_falls_back_to_uniform_legal_policy() -> None:
|
||||||
|
policy = regret_matching(
|
||||||
|
np.asarray([-1.0, 0.0, 3.0, 5.0], dtype=np.float32),
|
||||||
|
np.asarray([True, True, False, False]),
|
||||||
|
)
|
||||||
|
no_legal = regret_matching(
|
||||||
|
np.asarray([1.0, 2.0], dtype=np.float32),
|
||||||
|
np.asarray([False, False]),
|
||||||
|
)
|
||||||
|
|
||||||
|
np.testing.assert_allclose(policy, [0.5, 0.5, 0.0, 0.0])
|
||||||
|
np.testing.assert_allclose(no_legal, [0.0, 0.0])
|
||||||
|
|
||||||
|
|
||||||
|
def test_normalize_legal_policy_clamps_and_normalizes_legal_weights() -> None:
|
||||||
|
policy = normalize_legal_policy(
|
||||||
|
np.asarray([2.0, -1.0, 4.0, 6.0], dtype=np.float32),
|
||||||
|
np.asarray([True, True, False, True]),
|
||||||
|
)
|
||||||
|
fallback = normalize_legal_policy(
|
||||||
|
np.asarray([0.0, -1.0, 3.0], dtype=np.float32),
|
||||||
|
np.asarray([True, True, False]),
|
||||||
|
)
|
||||||
|
|
||||||
|
np.testing.assert_allclose(policy, [0.25, 0.0, 0.0, 0.75])
|
||||||
|
np.testing.assert_allclose(fallback, [0.5, 0.5, 0.0])
|
||||||
|
|
||||||
|
|
||||||
|
def test_sample_policy_uses_cumulative_probability_boundaries() -> None:
|
||||||
|
policy = np.asarray([0.2, 0.3, 0.5], dtype=np.float32)
|
||||||
|
|
||||||
|
assert sample_policy(policy, 0.0) == 0
|
||||||
|
assert sample_policy(policy, 0.21) == 1
|
||||||
|
assert sample_policy(policy, 0.51) == 2
|
||||||
|
assert sample_policy(policy, 0.999) == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_encode_info_state_is_deterministic_and_matches_legal_mask_tail() -> None:
|
||||||
|
state = GameState.new_game(LostCitiesConfig(seed=13))
|
||||||
|
|
||||||
|
encoded = encode_info_state(state, 0)
|
||||||
|
encoded_again = encode_info_state(state, 0)
|
||||||
|
legal_mask = np.asarray(state.unified_legal_mask(), dtype=np.float32)
|
||||||
|
|
||||||
|
assert encoded.dtype == np.float32
|
||||||
|
assert encoded.shape == (input_dim(state),)
|
||||||
|
np.testing.assert_array_equal(encoded, encoded_again)
|
||||||
|
np.testing.assert_array_equal(encoded[-len(legal_mask) :], legal_mask)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="invalid player"):
|
||||||
|
encode_info_state(state, 2)
|
||||||
Reference in New Issue
Block a user