Deep CFR advantage memory를 player별로 분리
This commit is contained in:
@@ -150,7 +150,9 @@ class DeepCFRTrainer:
|
||||
lr=self.config.optimization.learning_rate,
|
||||
weight_decay=self.config.optimization.weight_decay,
|
||||
)
|
||||
self.advantage_memory = ReservoirMemory(self.config.memory.advantage_capacity)
|
||||
self.advantage_memories = [
|
||||
ReservoirMemory(self.config.memory.advantage_capacity) for _ in range(2)
|
||||
]
|
||||
self.strategy_memory = ReservoirMemory(self.config.memory.strategy_capacity)
|
||||
self.rng = np.random.default_rng(self.config.run.seed + 101)
|
||||
self.iteration = 0
|
||||
@@ -217,6 +219,13 @@ class DeepCFRTrainer:
|
||||
self.strategy_optimizer.load_state_dict(payload["strategy_optimizer"])
|
||||
self.self_play_league_snapshots = payload.get("self_play_league_snapshots", [])
|
||||
|
||||
def _advantage_memory_size(self) -> int:
|
||||
return sum(len(memory) for memory in self.advantage_memories)
|
||||
|
||||
def _add_advantage_samples(self, samples: list[TrainingSample]) -> None:
|
||||
for sample in samples:
|
||||
self.advantage_memories[sample.player].add(sample, self.rng)
|
||||
|
||||
def run_iteration(self, iteration: int) -> IterationMetrics:
|
||||
self.iteration = iteration
|
||||
self._runtime_metrics = {}
|
||||
@@ -238,11 +247,13 @@ class DeepCFRTrainer:
|
||||
eval_started = time.perf_counter()
|
||||
eval_metrics = self._evaluate(iteration)
|
||||
self._runtime_metrics["evaluation_seconds"] = time.perf_counter() - eval_started
|
||||
self._runtime_metrics["advantage_memory_size"] = len(self.advantage_memory)
|
||||
self._runtime_metrics["advantage_memory_size"] = self._advantage_memory_size()
|
||||
for player, memory in enumerate(self.advantage_memories):
|
||||
self._runtime_metrics[f"advantage_player_{player}_memory_size"] = len(memory)
|
||||
self._runtime_metrics["strategy_memory_size"] = len(self.strategy_memory)
|
||||
return IterationMetrics(
|
||||
iteration=iteration,
|
||||
advantage_samples=len(self.advantage_memory),
|
||||
advantage_samples=self._advantage_memory_size(),
|
||||
strategy_samples=len(self.strategy_memory),
|
||||
advantage_loss=advantage_loss,
|
||||
strategy_loss=strategy_loss,
|
||||
@@ -313,7 +324,7 @@ class DeepCFRTrainer:
|
||||
)
|
||||
total_stats.accumulate(stats)
|
||||
memory_add_started = time.perf_counter()
|
||||
self.advantage_memory.add_many(advantage_samples, self.rng)
|
||||
self._add_advantage_samples(advantage_samples)
|
||||
self.strategy_memory.add_many(strategy_samples, self.rng)
|
||||
self._runtime_metrics["memory_add_seconds"] = (
|
||||
float(self._runtime_metrics.get("memory_add_seconds", 0.0))
|
||||
@@ -363,7 +374,7 @@ class DeepCFRTrainer:
|
||||
result = future.result()
|
||||
total_stats.accumulate(result.stats)
|
||||
memory_add_started = time.perf_counter()
|
||||
self.advantage_memory.add_many(result.advantage_samples, self.rng)
|
||||
self._add_advantage_samples(result.advantage_samples)
|
||||
self.strategy_memory.add_many(result.strategy_samples, self.rng)
|
||||
self._runtime_metrics["memory_add_seconds"] = (
|
||||
float(self._runtime_metrics.get("memory_add_seconds", 0.0))
|
||||
@@ -546,24 +557,20 @@ class DeepCFRTrainer:
|
||||
|
||||
def _train_advantage_networks(self) -> float:
|
||||
losses: list[float] = []
|
||||
for player, network in enumerate(self.advantage_networks):
|
||||
filter_started = time.perf_counter()
|
||||
samples = [sample for sample in self.advantage_memory.all() if sample.player == player]
|
||||
self._runtime_metrics[f"advantage_player_{player}_filter_seconds"] = (
|
||||
time.perf_counter() - filter_started
|
||||
)
|
||||
self._runtime_metrics[f"advantage_player_{player}_sample_count"] = len(samples)
|
||||
if not samples:
|
||||
for player, (network, memory) in enumerate(
|
||||
zip(self.advantage_networks, self.advantage_memories, strict=True)
|
||||
):
|
||||
self._runtime_metrics[f"advantage_player_{player}_sample_count"] = len(memory)
|
||||
if len(memory) == 0:
|
||||
continue
|
||||
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:
|
||||
samples = self.strategy_memory.all()
|
||||
self._runtime_metrics["strategy_sample_count"] = len(samples)
|
||||
if not samples:
|
||||
self._runtime_metrics["strategy_sample_count"] = len(self.strategy_memory)
|
||||
if len(self.strategy_memory) == 0:
|
||||
return 0.0
|
||||
return self._train_strategy(self.strategy_network, self.strategy_optimizer, samples)
|
||||
return self._train_strategy(self.strategy_network, self.strategy_optimizer)
|
||||
|
||||
def _batch_tensors(
|
||||
self,
|
||||
@@ -602,10 +609,9 @@ class DeepCFRTrainer:
|
||||
network.train()
|
||||
for _step in range(self.config.optimization.resolved_advantage_train_steps()):
|
||||
sample_started = time.perf_counter()
|
||||
batch = self.advantage_memory.sample(
|
||||
batch = self.advantage_memories[player].sample(
|
||||
self.config.optimization.resolved_advantage_batch_size(),
|
||||
self.rng,
|
||||
player=player,
|
||||
)
|
||||
self._runtime_metrics[f"advantage_player_{player}_sample_seconds"] = (
|
||||
float(self._runtime_metrics.get(f"advantage_player_{player}_sample_seconds", 0.0))
|
||||
@@ -630,7 +636,6 @@ class DeepCFRTrainer:
|
||||
self,
|
||||
network: nn.Module,
|
||||
optimizer: torch.optim.Optimizer,
|
||||
samples: list[TrainingSample],
|
||||
) -> float:
|
||||
last_loss = 0.0
|
||||
network.train()
|
||||
|
||||
Reference in New Issue
Block a user