Skip to content

REINFORCE

lerax.algorithm.REINFORCE

Bases: AbstractAlgorithm[PolicyType, REINFORCEState[PolicyType]]

REINFORCE algorithm with value function baseline.

Uses Monte Carlo returns (GAE with lambda=1) and vanilla policy gradient without clipping or importance sampling.

Attributes:

Name Type Description
optimizer optax.GradientTransformation

The optimizer used for training.

gae_lambda float

Lambda parameter for GAE, fixed to 1.0 for Monte Carlo returns.

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.

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.

1
num_steps int

Number of steps to run for each environment per update.

512
gamma float

Discount factor.

0.99
normalize_advantages bool

Whether to normalize advantages.

True
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
learning_rate optax.ScalarOrSchedule

Learning rate for the optimizer.

0.0003

optimizer instance-attribute

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

gae_lambda instance-attribute

gae_lambda: float = 1.0

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

value_loss_coefficient instance-attribute

value_loss_coefficient: float = value_loss_coefficient

max_grad_norm instance-attribute

max_grad_norm: float = max_grad_norm

reinforce_loss_grad class-attribute instance-attribute

reinforce_loss_grad = staticmethod(
    eqx.filter_value_and_grad(reinforce_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 = 1,
    num_steps: int = 512,
    gamma: float = 0.99,
    normalize_advantages: bool = True,
    value_loss_coefficient: float = 0.5,
    max_grad_norm: float = 0.5,
    learning_rate: optax.ScalarOrSchedule = 0.0003,
)

num_iterations

num_iterations(total_timesteps: int) -> int

step

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

Perform a single environment step and collect rollout data.

collect_rollout

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

Collect a rollout using the current policy.

reset

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

iteration

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

reinforce_loss staticmethod

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

train

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