From 66d02f773e9ab9f1fd9dd025e344bdce85a322fe Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EC=A0=95=EC=8B=9C=EC=9B=90?= Date: Wed, 6 May 2026 23:45:50 +0900 Subject: [PATCH] =?UTF-8?q?Deep=20CFR=20reservoir=20memory=20=EC=B6=94?= =?UTF-8?q?=EA=B0=80?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../games/classic/deep_cfr/config.py | 2 + .../games/classic/deep_cfr/memory.py | 43 +++++++++++++++++-- .../games/classic/deep_cfr/trainer.py | 22 +++++----- .../games/classic/deep_cfr/traverser.py | 6 ++- tests/games/classic/test_deep_cfr_trainer.py | 23 ++++++++++ 5 files changed, 80 insertions(+), 16 deletions(-) diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/config.py b/src/coolrl_lost_cities/games/classic/deep_cfr/config.py index c2ddb68..3a8f59e 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/config.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/config.py @@ -20,6 +20,8 @@ class DeepCFRConfig: strategy_sample_interval: int = 1 store_strategy_on_traverser_nodes: bool = True store_strategy_on_opponent_nodes: bool = True + advantage_memory_capacity: int = 2_000_000 + strategy_memory_capacity: int = 2_000_000 advantage_train_steps: int = 1 strategy_train_steps: int = 1 batch_size: int = 32 diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/memory.py b/src/coolrl_lost_cities/games/classic/deep_cfr/memory.py index 567b692..0078f25 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/memory.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/memory.py @@ -18,14 +18,49 @@ class ReservoirMemory: def __init__(self, capacity: int | None = None) -> None: self.capacity = capacity self._samples: list[TrainingSample] = [] + self.seen = 0 def __len__(self) -> int: return len(self._samples) - def add(self, sample: TrainingSample) -> None: - self._samples.append(sample) - if self.capacity is not None and len(self._samples) > self.capacity: - del self._samples[0 : len(self._samples) - self.capacity] + def add(self, sample: TrainingSample, rng: np.random.Generator | None = None) -> None: + self.seen += 1 + sample = TrainingSample( + info_state=np.asarray(sample.info_state, dtype=np.float32).copy(), + target=np.asarray(sample.target, dtype=np.float32).copy(), + legal_mask=np.asarray(sample.legal_mask, dtype=bool).copy(), + iteration=int(sample.iteration), + player=int(sample.player), + ) + if self.capacity is None or len(self._samples) < self.capacity: + self._samples.append(sample) + return + rng = rng or np.random.default_rng() + index = int(rng.integers(0, self.seen)) + if index < self.capacity: + self._samples[index] = sample + + def extend(self, samples: list[TrainingSample], rng: np.random.Generator | None = None) -> None: + for sample in samples: + self.add(sample, rng) def all(self) -> list[TrainingSample]: return list(self._samples) + + def sample( + self, + batch_size: int, + rng: np.random.Generator, + *, + player: int | None = None, + ) -> list[TrainingSample]: + candidates = ( + self._samples + if player is None + else [sample for sample in self._samples if sample.player == player] + ) + if not candidates: + raise ValueError("cannot sample from empty memory") + size = min(int(batch_size), len(candidates)) + indices = rng.choice(len(candidates), size=size, replace=len(candidates) < size) + return [candidates[int(index)] for index in indices] diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py b/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py index c9d7c38..daf65db 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py @@ -59,8 +59,8 @@ class DeepCFRTrainer: self.strategy_optimizer = torch.optim.Adam( self.strategy_network.parameters(), lr=self.config.learning_rate ) - self.advantage_memory = ReservoirMemory() - self.strategy_memory = ReservoirMemory() + self.advantage_memory = ReservoirMemory(self.config.advantage_memory_capacity) + self.strategy_memory = ReservoirMemory(self.config.strategy_memory_capacity) self.rng = np.random.default_rng(self.config.seed + 101) def run_iteration(self, iteration: int) -> IterationMetrics: @@ -119,9 +119,7 @@ class DeepCFRTrainer: samples = [sample for sample in self.advantage_memory.all() if sample.player == player] if not samples: continue - losses.append( - self._train_advantage(network, self.advantage_optimizers[player], samples) - ) + losses.append(self._train_advantage(player, network, self.advantage_optimizers[player])) return float(np.mean(losses)) if losses else 0.0 def _train_strategy_network(self) -> float: @@ -161,14 +159,16 @@ class DeepCFRTrainer: def _train_advantage( self, + player: int, network: nn.Module, optimizer: torch.optim.Optimizer, - samples: list[TrainingSample], ) -> float: last_loss = 0.0 network.train() - for step in range(max(self.config.advantage_train_steps, 0)): - x, y, legal = self._batch_tensors(self._batch(samples, step)) + for _step in range(max(self.config.advantage_train_steps, 0)): + x, y, legal = self._batch_tensors( + self.advantage_memory.sample(self.config.batch_size, self.rng, player=player) + ) pred = network(x) diff = (pred - y).masked_fill(~legal, 0.0) loss = diff.square().sum() / legal.sum().clamp_min(1) @@ -186,8 +186,10 @@ class DeepCFRTrainer: ) -> float: last_loss = 0.0 network.train() - for step in range(max(self.config.strategy_train_steps, 0)): - x, y, legal = self._batch_tensors(self._batch(samples, step)) + for _step in range(max(self.config.strategy_train_steps, 0)): + x, y, legal = self._batch_tensors( + self.strategy_memory.sample(self.config.batch_size, self.rng) + ) logits = network(x).masked_fill(~legal, torch.finfo(torch.float32).min) log_probs = nn.functional.log_softmax(logits, dim=-1).masked_fill(~legal, 0.0) loss = -(y * log_probs).sum(dim=-1).mean() diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/traverser.py b/src/coolrl_lost_cities/games/classic/deep_cfr/traverser.py index 996fb9d..76f8bc6 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/traverser.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/traverser.py @@ -216,7 +216,8 @@ class DeepCFRTraverser: legal_mask=legal, iteration=iteration, player=player, - ) + ), + self.rng, ) stats.advantage_samples += 1 @@ -328,7 +329,8 @@ class DeepCFRTraverser: legal_mask=legal, iteration=iteration, player=player, - ) + ), + self.rng, ) stats.strategy_samples += 1 diff --git a/tests/games/classic/test_deep_cfr_trainer.py b/tests/games/classic/test_deep_cfr_trainer.py index 639b9cb..c4b6299 100644 --- a/tests/games/classic/test_deep_cfr_trainer.py +++ b/tests/games/classic/test_deep_cfr_trainer.py @@ -4,6 +4,7 @@ import numpy as np from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig from coolrl_lost_cities.games.classic.deep_cfr.config import DeepCFRConfig +from coolrl_lost_cities.games.classic.deep_cfr.memory import ReservoirMemory, TrainingSample from coolrl_lost_cities.games.classic.deep_cfr.trainer import DeepCFRTrainer from coolrl_lost_cities.games.classic.deep_cfr.traverser import DeepCFRTraverser @@ -126,3 +127,25 @@ def test_deep_cfr_traverser_supports_outcome_sampling_and_rollout_cutoffs() -> N unsampled_legal = sample.legal_mask.copy() unsampled_legal[np.nonzero(sample.target)[0]] = False assert np.all(sample.target[unsampled_legal] == 0.0) + + +def test_reservoir_memory_caps_samples_and_filters_player_batches() -> None: + memory = ReservoirMemory(capacity=3) + rng = np.random.default_rng(37) + for index in range(10): + memory.add( + TrainingSample( + info_state=np.asarray([index], dtype=np.float32), + target=np.asarray([index], dtype=np.float32), + legal_mask=np.asarray([True]), + iteration=index, + player=index % 2, + ), + rng, + ) + + assert len(memory) == 3 + assert memory.seen == 10 + player_one = memory.sample(8, rng, player=1) + assert player_one + assert all(sample.player == 1 for sample in player_one)