Skip to content

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 log(num_actions).

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 q_lr when None.

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

optimizer: optax.GradientTransformation = (
    optax.inject_hyperparams(optax.adam)(policy_lr)
)

q_optimizer class-attribute instance-attribute

q_optimizer: optax.GradientTransformation = (
    optax.inject_hyperparams(optax.adam)(q_lr)
)

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
    )
)

buffer_size instance-attribute

buffer_size: int = buffer_size

gamma instance-attribute

gamma: float = gamma

learning_starts instance-attribute

learning_starts: int = learning_starts

num_envs instance-attribute

num_envs: int = num_envs

num_steps instance-attribute

num_steps: int = num_steps

batch_size instance-attribute

batch_size: int = batch_size

tau instance-attribute

tau: float = tau

autotune instance-attribute

autotune: bool = autotune

initial_alpha instance-attribute

initial_alpha: float = initial_alpha

target_entropy_scale instance-attribute

target_entropy_scale: float = target_entropy_scale

q_width_size instance-attribute

q_width_size: int = q_width_size

q_depth instance-attribute

q_depth: int = q_depth

q_loss_grad class-attribute instance-attribute

q_loss_grad = staticmethod(
    eqx.filter_value_and_grad(q_loss)
)

actor_loss_grad class-attribute instance-attribute

actor_loss_grad = staticmethod(
    eqx.filter_value_and_grad(actor_loss)
)

alpha_loss_grad class-attribute instance-attribute

alpha_loss_grad = staticmethod(
    eqx.filter_value_and_grad(alpha_loss)
)

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,
)

num_iterations

num_iterations(total_timesteps: int) -> int

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

per_iteration(
    state: SACDiscreteState[PolicyType],
) -> SACDiscreteState[PolicyType]

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.

iteration_count instance-attribute

iteration_count: Int[Array, '']

step_state instance-attribute

step_state: SACDiscreteStepState[PolicyType]

env instance-attribute

env: AbstractEnvLike

policy instance-attribute

policy: PolicyType

opt_state instance-attribute

opt_state: optax.OptState

callback_state instance-attribute

callback_state: AbstractCallbackState

qf1 instance-attribute

qf1: DiscreteQNetwork

qf2 instance-attribute

qf2: DiscreteQNetwork

qf1_target instance-attribute

qf1_target: DiscreteQNetwork

qf2_target instance-attribute

qf2_target: DiscreteQNetwork

q_opt_state instance-attribute

q_opt_state: optax.OptState

log_alpha instance-attribute

log_alpha: Float[Array, '']

alpha_opt_state instance-attribute

alpha_opt_state: optax.OptState

target_entropy instance-attribute

target_entropy: Float[Array, '']

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

with_callback_states[A: AbstractAlgorithmState](
    callback_state: AbstractCallbackState,
) -> A

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.