SAC-Discrete
lerax.algorithm.SACDiscrete
Bases: AbstractAlgorithm[PolicyType, SACDiscreteState[PolicyType]]
Soft Actor-Critic for discrete action spaces (SAC-Discrete).
Uses twin Q-networks that output Q-values for all actions, a categorical policy, and entropy regularization over the discrete action distribution. Based on the SAC-Discrete paper by Christodoulou (2019).
Attributes:
| Name | Type | Description |
|---|---|---|
optimizer |
optax.GradientTransformation
|
The actor optimizer. |
q_optimizer |
optax.GradientTransformation
|
The Q-network optimizer. |
alpha_optimizer |
optax.GradientTransformation
|
The entropy coefficient optimizer. |
buffer_size |
int
|
The size of the replay buffer. |
gamma |
float
|
Discount factor for future rewards. |
learning_starts |
int
|
Number of initial steps to collect before training. |
num_envs |
int
|
Number of parallel environments. |
num_steps |
int
|
Number of steps per iteration. |
batch_size |
int
|
Batch size for training. |
tau |
float
|
Soft update coefficient for target networks. |
autotune |
bool
|
Whether to automatically tune the entropy coefficient. |
initial_alpha |
float
|
Initial entropy coefficient value. |
target_entropy_scale |
float
|
Scale for target entropy relative to
|
q_width_size |
int
|
Width of Q-network hidden layers. |
q_depth |
int
|
Depth of Q-network hidden layers. |
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
buffer_size
|
int
|
The size of the replay buffer. |
1000000
|
gamma
|
float
|
Discount factor for future rewards. |
0.99
|
learning_starts
|
int
|
Number of initial steps before training. |
20000
|
num_envs
|
int
|
Number of parallel environments. |
1
|
num_steps
|
int
|
Number of steps per iteration. |
4
|
batch_size
|
int
|
Batch size for training. |
64
|
tau
|
float
|
Soft update coefficient for target networks. |
1.0
|
autotune
|
bool
|
Whether to automatically tune entropy coefficient. |
True
|
initial_alpha
|
float
|
Initial entropy coefficient value. |
0.2
|
target_entropy_scale
|
float
|
Scale for target entropy. |
0.89
|
policy_lr
|
optax.ScalarOrSchedule
|
Learning rate for the actor. |
0.0003
|
q_lr
|
optax.ScalarOrSchedule
|
Learning rate for Q-networks. |
0.0003
|
alpha_lr
|
optax.ScalarOrSchedule | None
|
Learning rate for the entropy coefficient.
Defaults to |
None
|
q_width_size
|
int
|
Width of Q-network hidden layers. |
256
|
q_depth
|
int
|
Depth of Q-network hidden layers. |
2
|
optimizer
instance-attribute
q_optimizer
class-attribute
instance-attribute
alpha_optimizer
class-attribute
instance-attribute
alpha_optimizer: optax.GradientTransformation = (
optax.inject_hyperparams(optax.adam)(
alpha_lr if alpha_lr is not None else q_lr
)
)
q_loss_grad
class-attribute
instance-attribute
actor_loss_grad
class-attribute
instance-attribute
alpha_loss_grad
class-attribute
instance-attribute
consolidate_callbacks
consolidate_callbacks(
callback: Sequence[AbstractCallback]
| AbstractCallback
| None = None,
) -> AbstractCallback
learn
learn(
env: AbstractEnvLike,
policy: PolicyType,
total_timesteps: int,
*,
key: Key[Array, ""],
callback: Sequence[AbstractCallback]
| AbstractCallback
| None = None,
) -> PolicyType
Train the policy on the environment for a given number of timesteps.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
env
|
AbstractEnvLike
|
The environment to train on. |
required |
policy
|
PolicyType
|
The policy to train. |
required |
total_timesteps
|
int
|
The total number of timesteps to train for. |
required |
key
|
Key[Array, '']
|
A JAX PRNG key. |
required |
callback
|
Sequence[AbstractCallback] | AbstractCallback | None
|
A callback or list of callbacks to use during training. |
None
|
Returns:
| Name | Type | Description |
|---|---|---|
policy |
PolicyType
|
The trained policy. |
__init__
__init__(
*,
buffer_size: int = 1000000,
gamma: float = 0.99,
learning_starts: int = 20000,
num_envs: int = 1,
num_steps: int = 4,
batch_size: int = 64,
tau: float = 1.0,
autotune: bool = True,
initial_alpha: float = 0.2,
target_entropy_scale: float = 0.89,
policy_lr: optax.ScalarOrSchedule = 0.0003,
q_lr: optax.ScalarOrSchedule = 0.0003,
alpha_lr: optax.ScalarOrSchedule | None = None,
q_width_size: int = 256,
q_depth: int = 2,
)
step
step(
env: AbstractEnvLike,
policy: PolicyType,
state: SACDiscreteStepState[PolicyType],
*,
key: Key[Array, ""],
callback: AbstractCallback,
) -> SACDiscreteStepState[PolicyType]
Perform a single environment step and store in replay buffer.
collect_learning_starts
collect_learning_starts(
env: AbstractEnvLike,
policy: PolicyType,
step_state: SACDiscreteStepState[PolicyType],
callback: AbstractCallback,
key: Key[Array, ""],
) -> SACDiscreteStepState[PolicyType]
Collect random initial experience before training begins.
collect_rollout
collect_rollout(
env: AbstractEnvLike,
policy: PolicyType,
step_state: SACDiscreteStepState[PolicyType],
callback: AbstractCallback,
key: Key[Array, ""],
) -> SACDiscreteStepState[PolicyType]
Collect a rollout of experience into the replay buffer.
reset
reset(
env: AbstractEnvLike,
policy: PolicyType,
*,
key: Key[Array, ""],
callback: AbstractCallback,
) -> SACDiscreteState[PolicyType]
iteration
iteration(
state: SACDiscreteState[PolicyType],
*,
key: Key[Array, ""],
callback: AbstractCallback,
) -> SACDiscreteState[PolicyType]
per_iteration
Polyak-average target Q-networks.
q_loss
staticmethod
q_loss(
q_params: tuple[DiscreteQNetwork, DiscreteQNetwork],
batch: ReplayBuffer,
target: Float[Array, " batch_size"],
) -> Float[Array, ""]
Compute combined MSE loss for twin discrete Q-networks.
actor_loss
staticmethod
actor_loss(
policy: PolicyType,
batch: ReplayBuffer,
qf1: DiscreteQNetwork,
qf2: DiscreteQNetwork,
alpha: Float[Array, ""],
) -> Float[Array, ""]
Compute actor loss over discrete action distribution.
alpha_loss
staticmethod
alpha_loss(
log_alpha: Float[Array, ""],
action_probs: Float[Array, "batch_size num_actions"],
log_probs: Float[Array, "batch_size num_actions"],
target_entropy: Float[Array, ""],
) -> Float[Array, ""]
Compute entropy coefficient loss for discrete actions.
sac_discrete_train
sac_discrete_train(
policy: PolicyType,
opt_state: optax.OptState,
buffer: ReplayBuffer,
qf1: DiscreteQNetwork,
qf2: DiscreteQNetwork,
qf1_target: DiscreteQNetwork,
qf2_target: DiscreteQNetwork,
q_opt_state: optax.OptState,
log_alpha: Float[Array, ""],
alpha_opt_state: optax.OptState,
target_entropy: Float[Array, ""],
*,
key: Key[Array, ""],
) -> tuple[
PolicyType,
optax.OptState,
DiscreteQNetwork,
DiscreteQNetwork,
optax.OptState,
Float[Array, ""],
optax.OptState,
dict[str, Scalar],
]
lerax.algorithm.SACDiscreteState
Bases: AbstractAlgorithmState[PolicyType]
Iteration-level state for SAC-Discrete.
Attributes:
| Name | Type | Description |
|---|---|---|
iteration_count |
Int[Array, '']
|
The current iteration count. |
step_state |
SACDiscreteStepState[PolicyType]
|
The step-level state. |
env |
AbstractEnvLike
|
The environment being used. |
policy |
PolicyType
|
The policy being trained. |
opt_state |
optax.OptState
|
The actor optimizer state. |
callback_state |
AbstractCallbackState
|
The callback state. |
qf1 |
DiscreteQNetwork
|
First discrete Q-network. |
qf2 |
DiscreteQNetwork
|
Second discrete Q-network. |
qf1_target |
DiscreteQNetwork
|
Target for first Q-network. |
qf2_target |
DiscreteQNetwork
|
Target for second Q-network. |
q_opt_state |
optax.OptState
|
Optimizer state for Q-networks. |
log_alpha |
Float[Array, '']
|
Log of the entropy coefficient. |
alpha_opt_state |
optax.OptState
|
Optimizer state for entropy coefficient. |
target_entropy |
Float[Array, '']
|
Target entropy value. |
next
next[A: AbstractAlgorithmState](
step_state: AbstractStepState,
policy: PolicyType,
opt_state: optax.OptState,
) -> A
Return a new algorithm state for the next iteration.
Increments the iteration count and updates the step state, policy, and optimizer state.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
step_state
|
AbstractStepState
|
The new step state. |
required |
policy
|
PolicyType
|
The new policy. |
required |
opt_state
|
optax.OptState
|
The new optimizer state. |
required |
Returns:
| Type | Description |
|---|---|
A
|
A new algorithm state with the updated fields. |
with_callback_states
Return a new algorithm state with the given callback state.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
callback_state
|
AbstractCallbackState
|
The new callback state. |
required |
Returns:
| Type | Description |
|---|---|
A
|
A new algorithm state with the updated callback state. |