Deep CFR playability encoding 추가
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
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
|
||||
|
||||
from coolrl_lost_cities.games.classic.deep_cfr.benchmark import (
|
||||
@@ -53,6 +54,47 @@ def test_deep_cfr_loads_mapped_legacy_reproduction_config() -> None:
|
||||
assert config.checkpoint.save_iteration_interval == 10
|
||||
|
||||
|
||||
def test_deep_cfr_playability_encoding_extends_input_shape() -> None:
|
||||
state = GameState.new_game(LostCitiesConfig(seed=61), seed=61)
|
||||
base_dim = input_dim(state)
|
||||
derived_config = _deep_cfr_config({"encoding": {"derived_playability": True}})
|
||||
slot_config = _deep_cfr_config(
|
||||
{"encoding": {"derived_playability": True, "slot_aware_playability": True}}
|
||||
)
|
||||
|
||||
derived_dim = input_dim(state, derived_config.encoding)
|
||||
slot_dim = input_dim(state, slot_config.encoding)
|
||||
|
||||
assert derived_dim == base_dim + state.config.n_colors * 19 + 3
|
||||
assert slot_dim == derived_dim + state.config.hand_size * 12
|
||||
assert encode_info_state(state, 0, slot_config.encoding).shape == (slot_dim,)
|
||||
|
||||
|
||||
def test_deep_cfr_trainer_uses_playability_encoding() -> None:
|
||||
config = _deep_cfr_config(
|
||||
{
|
||||
"run": {"iterations": 1, "seed": 62},
|
||||
"encoding": {"derived_playability": True, "slot_aware_playability": True},
|
||||
"network": {"hidden_size": 16},
|
||||
"traversal": {"traversals_per_iteration": 1, "max_depth": 1, "max_nodes": 16},
|
||||
"optimization": {
|
||||
"advantage_train_steps": 1,
|
||||
"strategy_train_steps": 1,
|
||||
"batch_size": 2,
|
||||
},
|
||||
"checkpoint": {"save_every_iteration": False},
|
||||
}
|
||||
)
|
||||
game_config = LostCitiesConfig(seed=62)
|
||||
trainer = DeepCFRTrainer(config, game_config)
|
||||
|
||||
metrics = trainer.train()
|
||||
|
||||
probe = GameState.new_game(game_config, seed=62)
|
||||
assert trainer.input_dim == input_dim(probe, config.encoding)
|
||||
assert metrics[0].advantage_samples > 0
|
||||
|
||||
|
||||
def test_deep_cfr_trainer_smoke_run() -> None:
|
||||
trainer = DeepCFRTrainer(
|
||||
_deep_cfr_config(
|
||||
|
||||
Reference in New Issue
Block a user