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