From b5b4f97d414428d656bcd4395dbdfc19821403b3 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EC=A0=95=EC=8B=9C=EC=9B=90?= Date: Thu, 7 May 2026 02:32:24 +0900 Subject: [PATCH] =?UTF-8?q?Deep=20CFR=20slot-aware=20encoding=20NaN=20?= =?UTF-8?q?=EC=88=98=EC=A0=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../games/classic/deep_cfr/encoding.pyx | 2 ++ tests/games/classic/test_deep_cfr_trainer.py | 14 ++++++++++++++ 2 files changed, 16 insertions(+) diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/encoding.pyx b/src/coolrl_lost_cities/games/classic/deep_cfr/encoding.pyx index 6fcd454..062a567 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/encoding.pyx +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/encoding.pyx @@ -236,6 +236,8 @@ cdef int _append_slot_aware_playability_features_c(GameState state, int player, for slot in range(state.hand_size): if slot >= state.hand_lens[player]: + for color in range(SLOT_AWARE_PLAYABILITY_PER_SLOT): + out[idx + color] = 0.0 idx += SLOT_AWARE_PLAYABILITY_PER_SLOT continue card = state.hand_cards[state._hand_index(player, slot)] diff --git a/tests/games/classic/test_deep_cfr_trainer.py b/tests/games/classic/test_deep_cfr_trainer.py index 97cd192..241d57d 100644 --- a/tests/games/classic/test_deep_cfr_trainer.py +++ b/tests/games/classic/test_deep_cfr_trainer.py @@ -165,6 +165,20 @@ def test_deep_cfr_playability_encoding_extends_input_shape() -> None: assert encode_info_state(state, 0, slot_config.encoding).shape == (slot_dim,) +def test_deep_cfr_slot_aware_playability_encoding_zero_fills_empty_hand_slots() -> None: + config = _deep_cfr_config( + {"encoding": {"derived_playability": True, "slot_aware_playability": True}} + ) + state = GameState.new_game(LostCitiesConfig(seed=63), seed=63) + state.phase = "draw" + state.pending_discarded_color = -1 + state.apply_action(0) + + encoded = encode_info_state(state, 0, config.encoding) + + assert np.isfinite(encoded).all() + + def test_deep_cfr_trainer_uses_playability_encoding() -> None: config = _deep_cfr_config( {