mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-24 10:16:09 +00:00
refactor(sac): simplify optimizer return structure
This commit is contained in:
@@ -457,14 +457,10 @@ class SACAlgorithm(RLAlgorithm):
|
|||||||
policy (nn.Module): The policy model containing the actor, critic, and temperature components.
|
policy (nn.Module): The policy model containing the actor, critic, and temperature components.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Tuple[Dict[str, torch.optim.Optimizer], Optional[torch.optim.lr_scheduler._LRScheduler]]:
|
A dictionary mapping component names ("actor", "critic", "temperature")
|
||||||
A tuple containing:
|
to their respective Adam optimizers.
|
||||||
- `optimizers`: A dictionary mapping component names ("actor", "critic", "temperature") to their respective Adam optimizers.
|
|
||||||
- `lr_scheduler`: Currently set to `None` but can be extended to support learning rate scheduling.
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
actor_params = self.policy.get_optim_params()["actor"]
|
actor_params = self.policy.get_optim_params()["actor"]
|
||||||
lr_scheduler = None
|
|
||||||
self.optimizers = {
|
self.optimizers = {
|
||||||
"actor": torch.optim.Adam(actor_params, lr=self.config.actor_lr),
|
"actor": torch.optim.Adam(actor_params, lr=self.config.actor_lr),
|
||||||
"critic": torch.optim.Adam(self.critic_ensemble.parameters(), lr=self.config.critic_lr),
|
"critic": torch.optim.Adam(self.critic_ensemble.parameters(), lr=self.config.critic_lr),
|
||||||
@@ -474,7 +470,7 @@ class SACAlgorithm(RLAlgorithm):
|
|||||||
self.optimizers["discrete_critic"] = torch.optim.Adam(
|
self.optimizers["discrete_critic"] = torch.optim.Adam(
|
||||||
self.discrete_critic.parameters(), lr=self.config.critic_lr
|
self.discrete_critic.parameters(), lr=self.config.critic_lr
|
||||||
)
|
)
|
||||||
return self.optimizers, lr_scheduler
|
return self.optimizers
|
||||||
|
|
||||||
def get_optimizers(self) -> dict[str, Optimizer]:
|
def get_optimizers(self) -> dict[str, Optimizer]:
|
||||||
return self.optimizers
|
return self.optimizers
|
||||||
|
|||||||
Reference in New Issue
Block a user