Deep CFR legacy 학습 dynamics 정렬
This commit is contained in:
@@ -181,14 +181,14 @@ class DeepCFRTrainer:
|
||||
self.input_dim, self.action_size, self.config.network
|
||||
).to(self.device)
|
||||
self.advantage_optimizers = [
|
||||
torch.optim.Adam(
|
||||
torch.optim.AdamW(
|
||||
network.parameters(),
|
||||
lr=self.config.optimization.learning_rate,
|
||||
weight_decay=self.config.optimization.weight_decay,
|
||||
)
|
||||
for network in self.advantage_networks
|
||||
]
|
||||
self.strategy_optimizer = torch.optim.Adam(
|
||||
self.strategy_optimizer = torch.optim.AdamW(
|
||||
self.strategy_network.parameters(),
|
||||
lr=self.config.optimization.learning_rate,
|
||||
weight_decay=self.config.optimization.weight_decay,
|
||||
@@ -717,7 +717,7 @@ class DeepCFRTrainer:
|
||||
network: nn.Module,
|
||||
optimizer: torch.optim.Optimizer,
|
||||
) -> float:
|
||||
last_loss = 0.0
|
||||
losses: list[float] = []
|
||||
network.train()
|
||||
for _step in range(self.config.optimization.resolved_advantage_train_steps()):
|
||||
sample_started = time.perf_counter()
|
||||
@@ -741,8 +741,8 @@ class DeepCFRTrainer:
|
||||
network.parameters(), self.config.optimization.grad_clip
|
||||
)
|
||||
optimizer.step()
|
||||
last_loss = float(loss.detach().cpu())
|
||||
return last_loss
|
||||
losses.append(float(loss.detach().cpu()))
|
||||
return float(np.mean(losses)) if losses else 0.0
|
||||
|
||||
def _train_strategy(
|
||||
self,
|
||||
|
||||
@@ -94,6 +94,8 @@ cdef class CythonDeepCFRTraverser:
|
||||
cdef float self_play_older_weight
|
||||
cdef float self_play_anchor_weight
|
||||
cdef int self_play_recent_window
|
||||
cdef int active_self_play_bucket
|
||||
cdef object active_self_play_networks
|
||||
cdef int endpoint_depth_bucket_width
|
||||
cdef int endpoint_depth_bucket_max
|
||||
cdef bint derived_playability
|
||||
@@ -187,6 +189,8 @@ cdef class CythonDeepCFRTraverser:
|
||||
self.self_play_older_weight = max(0.0, self_play_older_weight)
|
||||
self.self_play_anchor_weight = max(0.0, self_play_anchor_weight)
|
||||
self.self_play_recent_window = max(0, self_play_recent_window)
|
||||
self.active_self_play_bucket = 0
|
||||
self.active_self_play_networks = None
|
||||
self.safe_heuristic_opponent_bot = (
|
||||
SafeHeuristicBot()
|
||||
if self.opponent_policy_id == 1 or self.self_play_anchor_probability > 0.0
|
||||
@@ -203,8 +207,20 @@ cdef class CythonDeepCFRTraverser:
|
||||
self.input_dim = _input_dim_with_flags_c(
|
||||
state, self.derived_playability, self.slot_aware_playability
|
||||
)
|
||||
value = self._traverse(state, traverser, iteration, 0, stats)
|
||||
return value, stats
|
||||
if self.opponent_policy_id == 2:
|
||||
self.active_self_play_bucket = self._self_play_bucket()
|
||||
if self.active_self_play_bucket in (1, 2):
|
||||
self.active_self_play_networks = self._self_play_snapshot_networks(
|
||||
self.active_self_play_bucket
|
||||
)
|
||||
else:
|
||||
self.active_self_play_networks = None
|
||||
try:
|
||||
value = self._traverse(state, traverser, iteration, 0, stats)
|
||||
return value, stats
|
||||
finally:
|
||||
self.active_self_play_bucket = 0
|
||||
self.active_self_play_networks = None
|
||||
|
||||
cdef float _traverse(
|
||||
self,
|
||||
@@ -401,14 +417,14 @@ cdef class CythonDeepCFRTraverser:
|
||||
if self.safe_heuristic_opponent_bot is None:
|
||||
self.safe_heuristic_opponent_bot = SafeHeuristicBot()
|
||||
return int(self.safe_heuristic_opponent_bot.act(state))
|
||||
bucket = self._self_play_bucket()
|
||||
bucket = self.active_self_play_bucket
|
||||
if bucket == 0:
|
||||
return -1
|
||||
if bucket == 3:
|
||||
if self.safe_heuristic_opponent_bot is None:
|
||||
self.safe_heuristic_opponent_bot = SafeHeuristicBot()
|
||||
return int(self.safe_heuristic_opponent_bot.act(state))
|
||||
networks = self._self_play_snapshot_networks(bucket)
|
||||
networks = self.active_self_play_networks
|
||||
if networks is None:
|
||||
return -1
|
||||
self._policy_from_networks(networks, state, player, legal, policy)
|
||||
|
||||
Reference in New Issue
Block a user