Implementing Custom Algorithms
Implementing custom algorithms is only supported with the fsdp and megatron backends
SkyRL provides a registry system for easily implementing custom algorithms (advantage estimators, policy loss) without modifying the core codebase.
The API for the registry system can be found in skyrl/backends/skyrl_train/utils/ppo_utils.py.
Example scripts of using the registry can be found at examples/train/algorithms/.
Additionally for more control, you can subclass the BasePPOExp class from skyrl/train/entrypoints/main_base.py and override the BasePPOExp.get_trainer method to return a custom trainer class.
This allows you to have full control over the training loop and implementing custom reward functions and output postprocessing.
We provide an example of this for applying custom reward penalties in our DAPO example.
Registering a Custom Advantage Estimator
You can register custom advantage estimators using either a decorator or the registry directly:
from skyrl.backends.skyrl_train.utils.ppo_utils import register_advantage_estimator, AdvantageEstimatorRegistry
import torch
# Using the decorator
@register_advantage_estimator("simple_baseline")
def compute_simple_baseline_advantage(
token_level_rewards: torch.Tensor, response_mask: torch.Tensor, index: np.ndarray, **kwargs
):
with torch.no_grad():
response_rewards = (token_level_rewards * response_mask).sum(dim=-1, keepdim=True)
# Simple baseline: use the mean reward across the batch
baseline = response_rewards.mean()
advantages = (response_rewards - baseline) * response_mask
returns = advantages.clone()
return advantages, returns
# Or register directly
def another_estimator(**kwargs):
# Implementation here
pass
AdvantageEstimatorRegistry.register("direct_registration", another_estimator)Registering a Custom Policy Loss
Similarly, you can register custom policy loss functions:
from skyrl.backends.skyrl_train.utils.ppo_utils import register_policy_loss, PolicyLossRegistry
@register_policy_loss("reinforce")
def compute_reinforce_policy_loss(log_probs, old_log_probs, advantages, config, loss_mask=None, rollout_log_probs=None):
# Your custom policy loss implementation (like REINFORCE)
loss = (-log_probs * advantages).mean()
# return loss and loss metrics
return loss, {"clip_ratio": 0.0}How the built-in PPO loss composes the config
It can be helpful to see how the built-in loss uses the trainer.algorithm options before writing
your own. ppo_policy_loss
(which handles policy_loss_type of regular and dual_clip) applies them in this order:
- Computes the importance ratio
exp(log_probs - old_log_probs). - Clips it to
[1 - eps_clip_low, 1 + eps_clip_high]and takes the pessimistic (minimum) of the clipped and unclipped surrogate objectives. - For
dual_clip, additionally caps the loss at-advantages * clip_ratio_con the tokens where the advantage is negative. - Applies
off_policy_correction, which can both rescale the loss and mask out tokens or sequences. - Sums the masked per-token loss via
reduce_loss.
Note that step 5 is a plain masked sum — loss_reduction is not applied by the loss function.
The trainer pre-scales advantages per mini-batch (apply_loss_reduction_to_advantages_minibatch)
before the loss is called, so the reduction is already baked into the advantages. A custom policy
loss should therefore sum its masked per-token losses rather than averaging them, or it will apply
the reduction twice.
Custom losses receive the same config object (AlgorithmConfig), so any field on it — including
fields added by subclassing — is available. See the
configuration API reference for the full field list.
Registry Ray Distribution
The registry system handles Ray actor synchronization when Ray is initialized. Functions registered on one process will be available to all Ray actors:
import ray
from skyrl.backends.skyrl_train.utils.ppo_utils import AdvantageEstimatorRegistry, sync_registries
# Register a function on the main process
def my_function(**kwargs):
# A dummy function for demonstration
pass
AdvantageEstimatorRegistry.register("my_function", my_function)
# After Ray is initialized, we sync the registries to a named ray actor (in utils/utils.py::initialize_ray)
ray.init()
sync_registries()
@ray.remote(num_cpus=1)
def skyrl_entrypoint(cfg: SkyRLTrainConfig):
# Function is now available on all Ray processes
available_functions = AdvantageEstimatorRegistry.list_available() # will include "my_function"
exp = BasePPOExp(cfg)
exp.run()Creating a Custom Trainer
To create a custom trainer for full control of your training loop, you can subclass the BasePPOExp class from skyrl/train/entrypoints/main_base.py and override the BasePPOExp.get_trainer method to return a custom trainer class.
We show the outline of creating a custom trainer below, and you can find a full running example in our DAPO example.
class CustomTrainer(RayPPOTrainer):
@torch.no_grad()
def postprocess_generator_output(
self, generator_output: GeneratorOutput, uids: List[str]
) -> Tuple[GeneratorOutput, List[str]]:
# apply custom reward penalties
...
# use base class impl for metrics and per-token reward conversion
return super().postprocess_generator_output(generator_output, uids)
class CustomExp(BasePPOExp):
def get_trainer(self, *args, **kwargs):
return CustomTrainer(*args, **kwargs)
@ray.remote(num_cpus=1)
def skyrl_entrypoint(cfg: SkyRLTrainConfig):
exp = CustomExp(cfg)
exp.run()
...
if __name__ == "__main__":
main()