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
|
reinforce_loss_grad
class-attribute
instance-attribute
per_iteration
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,
)
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]