Skip to content

A2C

lerax.algorithm.A2C

Bases: AbstractAlgorithm[PolicyType, A2CState[PolicyType]]

Advantage Actor-Critic (A2C) algorithm.

Uses GAE for advantage estimation and performs a single gradient update per rollout with entropy regularization.

Attributes:

Name Type Description
optimizer optax.GradientTransformation

The optimizer used for training.

gae_lambda float

Lambda parameter for Generalized Advantage Estimation (GAE).

gamma float

Discount factor.

num_envs int

Number of parallel environments.

num_steps int

Number of steps to run for each environment per update.

batch_size int

Size of each training batch.

normalize_advantages bool

Whether to normalize advantages.

entropy_loss_coefficient float

Coefficient for the entropy loss term.

value_loss_coefficient float

Coefficient for the value function loss term.

max_grad_norm float

Maximum gradient norm for gradient clipping.

Parameters:

Name Type Description Default
num_envs int

Number of parallel environments.

4
num_steps int

Number of steps to run for each environment per update.

5
gae_lambda float

Lambda parameter for Generalized Advantage Estimation (GAE).

1.0
gamma float

Discount factor.

0.99
entropy_loss_coefficient float

Coefficient for the entropy loss term.

0.0
value_loss_coefficient float

Coefficient for the value function loss term.

0.5
max_grad_norm float

Maximum gradient norm for gradient clipping.

0.5
normalize_advantages bool

Whether to normalize advantages.

False
learning_rate optax.ScalarOrSchedule

Learning rate for the optimizer.

0.0007

optimizer instance-attribute

optimizer: optax.GradientTransformation = optax.chain(
    clip, adam
)

gae_lambda instance-attribute

gae_lambda: float = gae_lambda

gamma instance-attribute

gamma: float = gamma

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 = self.num_steps * self.num_envs

normalize_advantages instance-attribute

normalize_advantages: bool = normalize_advantages

entropy_loss_coefficient instance-attribute

entropy_loss_coefficient: float = entropy_loss_coefficient

value_loss_coefficient instance-attribute

value_loss_coefficient: float = value_loss_coefficient

max_grad_norm instance-attribute

max_grad_norm: float = max_grad_norm

a2c_loss_grad class-attribute instance-attribute

a2c_loss_grad = staticmethod(
    eqx.filter_value_and_grad(a2c_loss, has_aux=True)
)

per_iteration

per_iteration(state: StateType) -> StateType

Process the algorithm state after each iteration.

Used for algorithm-specific bookkeeping. Default is identity. Override for target network updates, etc.

Parameters:

Name Type Description Default
state StateType

The current algorithm state.

required

Returns:

Type Description
StateType

The updated algorithm state.

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__(
    *,
    num_envs: int = 4,
    num_steps: int = 5,
    gae_lambda: float = 1.0,
    gamma: float = 0.99,
    entropy_loss_coefficient: float = 0.0,
    value_loss_coefficient: float = 0.5,
    max_grad_norm: float = 0.5,
    normalize_advantages: bool = False,
    learning_rate: optax.ScalarOrSchedule = 0.0007,
)

num_iterations

num_iterations(total_timesteps: int) -> int

step

step(
    env: AbstractEnvLike,
    policy: PolicyType,
    state: A2CStepState[PolicyType],
    *,
    key: Key[Array, ""],
    callback: AbstractCallback,
) -> tuple[A2CStepState[PolicyType], RolloutBuffer]

Perform a single environment step and collect rollout data.

collect_rollout

collect_rollout(
    env: AbstractEnvLike,
    policy: PolicyType,
    step_state: A2CStepState[PolicyType],
    callback: AbstractCallback,
    key: Key[Array, ""],
) -> tuple[A2CStepState[PolicyType], RolloutBuffer]

Collect a rollout using the current policy.

reset

reset(
    env: AbstractEnvLike,
    policy: PolicyType,
    *,
    key: Key[Array, ""],
    callback: AbstractCallback,
) -> A2CState[PolicyType]

iteration

iteration(
    state: A2CState[PolicyType],
    *,
    key: Key[Array, ""],
    callback: AbstractCallback,
) -> A2CState[PolicyType]

a2c_loss staticmethod

a2c_loss(
    policy: PolicyType,
    rollout_buffer: RolloutBuffer,
    normalize_advantages: bool,
    value_loss_coefficient: float,
    entropy_loss_coefficient: float,
) -> tuple[Float[Array, ""], A2CStats]

train

train(
    policy: PolicyType,
    opt_state: optax.OptState,
    buffer: RolloutBuffer,
    *,
    key: Key[Array, ""],
) -> tuple[PolicyType, optax.OptState, dict[str, Scalar]]