Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions docs/source/changelog.rst
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,15 @@ Changelog
Upcoming version (not yet released)
-----------------------------------

Added
^^^^^

- Added FlashSAC (off-policy Soft Actor-Critic) support. New RSL-RL config
dataclasses (``RslRlOffPolicyRunnerCfg`` and the ``RslRlFlashSac*`` /
``RslRlReplayBufferCfg`` configs) and ``MjlabOffPolicyRunner`` wire the
``FlashSAC`` algorithm and ``OffPolicyRunner`` into mjlab, selectable purely
by config. Registered the ``Mjlab-Velocity-Flat-Unitree-G1-FlashSAC`` task.

Version 1.6.0 (August 8, 2026)
------------------------------

Expand Down
6 changes: 6 additions & 0 deletions src/mjlab/rl/__init__.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,12 @@
from mjlab.rl.config import RslRlBaseRunnerCfg as RslRlBaseRunnerCfg
from mjlab.rl.config import RslRlFlashSacActorCfg as RslRlFlashSacActorCfg
from mjlab.rl.config import RslRlFlashSacAlgorithmCfg as RslRlFlashSacAlgorithmCfg
from mjlab.rl.config import RslRlFlashSacCriticCfg as RslRlFlashSacCriticCfg
from mjlab.rl.config import RslRlModelCfg as RslRlModelCfg
from mjlab.rl.config import RslRlOffPolicyRunnerCfg as RslRlOffPolicyRunnerCfg
from mjlab.rl.config import RslRlOnPolicyRunnerCfg as RslRlOnPolicyRunnerCfg
from mjlab.rl.config import RslRlPpoAlgorithmCfg as RslRlPpoAlgorithmCfg
from mjlab.rl.config import RslRlReplayBufferCfg as RslRlReplayBufferCfg
from mjlab.rl.runner import MjlabOffPolicyRunner as MjlabOffPolicyRunner
from mjlab.rl.runner import MjlabOnPolicyRunner as MjlabOnPolicyRunner
from mjlab.rl.vecenv_wrapper import RslRlVecEnvWrapper as RslRlVecEnvWrapper
135 changes: 135 additions & 0 deletions src/mjlab/rl/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -143,3 +143,138 @@ class RslRlOnPolicyRunnerCfg(RslRlBaseRunnerCfg):
"""The critic configuration."""
algorithm: RslRlPpoAlgorithmCfg = field(default_factory=RslRlPpoAlgorithmCfg)
"""The algorithm configuration."""


@dataclass
class RslRlFlashSacActorCfg:
"""Config for the FlashSAC actor model."""

num_blocks: int = 2
"""Number of residual blocks in the actor trunk."""
hidden_dim: int = 128
"""Hidden dimension of the actor."""
obs_normalization: bool = False
"""Whether to apply empirical observation normalization (the network also
self-normalizes via BatchNorm, so this is off by default)."""
log_std_min: float = -5.0
"""Lower bound for the (squashed) policy log-std."""
log_std_max: float = 2.0
"""Upper bound for the (squashed) policy log-std."""
class_name: str = "FlashSACActorModel"
"""Model class name resolved by RSL-RL."""


@dataclass
class RslRlFlashSacCriticCfg:
"""Config for the FlashSAC distributional double critic."""

num_blocks: int = 2
"""Number of residual blocks per critic ensemble member."""
hidden_dim: int = 256
"""Hidden dimension of the critic."""
num_bins: int = 101
"""Number of atoms in the categorical (C51) value distribution."""
min_v: float = -5.0
"""Lower bound of the value support. Should match ``-normalized_g_max``."""
max_v: float = 5.0
"""Upper bound of the value support. Should match ``normalized_g_max``."""
num_qs: int = 2
"""Number of Q-ensemble members (clipped double-Q uses 2)."""
obs_normalization: bool = False
"""Whether to apply empirical observation normalization."""
class_name: str = "FlashSACCriticModel"
"""Model class name resolved by RSL-RL."""


@dataclass
class RslRlReplayBufferCfg:
"""Config for the off-policy replay buffer."""

capacity: int = 1_000_000
"""Maximum number of transitions stored."""
min_length: int = 10_000
"""Minimum number of transitions before sampling/updates begin."""
sample_batch_size: int = 2048
"""Mini-batch size drawn from the buffer per gradient step."""


@dataclass
class RslRlFlashSacAlgorithmCfg:
"""Config for the FlashSAC algorithm."""

gamma: float = 0.99
"""The discount factor."""
n_step: int = 1
"""Number of steps for n-step return accumulation."""
learning_rate_init: float = 3e-4
"""Initial learning rate (start of warmup)."""
learning_rate_peak: float = 3e-4
"""Peak learning rate (end of warmup)."""
learning_rate_end: float = 1.5e-4
"""Final learning rate after cosine decay."""
learning_rate_warmup_steps: int = 0
"""Number of linear warmup steps."""
learning_rate_decay_steps: int = 1_000_000
"""Total schedule length (warmup + cosine decay), in gradient steps."""
critic_target_update_tau: float = 0.01
"""EMA coefficient for the target critic update."""
num_bins: int = 101
"""Number of atoms in the categorical TD target. Must match the critic."""
min_v: float = -5.0
"""Lower bound of the value support. Must match the critic and ``-normalized_g_max``."""
max_v: float = 5.0
"""Upper bound of the value support. Must match the critic and ``normalized_g_max``."""
temp_initial_value: float = 0.01
"""Initial entropy-temperature value."""
temp_target_sigma: float = 0.15
"""Target action std used to auto-compute the target entropy when
``temp_target_entropy`` is None."""
temp_target_entropy: float | None = None
"""Target entropy. If None, it is auto-computed from the action dim and
``temp_target_sigma``."""
actor_update_period: int = 2
"""Delayed policy update period (actor/temperature update every N critic updates)."""
actor_bc_alpha: float = 0.0
"""Behavior-cloning regularization coefficient (0 disables it)."""
actor_noise_zeta_mu: float = 2.0
"""Zeta-distribution exponent for action-noise repetition."""
actor_noise_zeta_max: int = 16
"""Maximum noise-repetition length."""
normalize_reward: bool = True
"""Whether to normalize rewards (required True in this version)."""
normalized_g_max: float = 5.0
"""Return-normalization cap; also sets the critic value support magnitude."""
use_amp: bool = False
"""Whether to use automatic mixed precision (must be False in this version)."""
class_name: str = "FlashSAC"
"""Algorithm class name resolved by RSL-RL."""


@dataclass
class RslRlOffPolicyRunnerCfg(RslRlBaseRunnerCfg):
"""Runner config for off-policy (FlashSAC) training.

A drop-in sibling of :class:`RslRlOnPolicyRunnerCfg`: it reuses the base
runner fields (``obs_groups``, ``num_steps_per_env``, ``save_interval``,
``clip_actions``, ...) and adds the FlashSAC model/algorithm/replay configs.
These dataclass defaults are the single source of default hyperparameters;
``asdict`` materializes them into the plain dict RSL-RL consumes (fail-loud:
RSL-RL substitutes no defaults of its own).
"""

class_name: str = "OffPolicyRunner"
"""The runner class name."""
updates_per_step: float = 1.0
"""Gradient updates per collected environment step (may be < 1.0)."""
torch_compile_mode: str | None = None
"""torch.compile mode. Must be None in this version (eager-only)."""
actor: RslRlFlashSacActorCfg = field(default_factory=RslRlFlashSacActorCfg)
"""The actor configuration."""
critic: RslRlFlashSacCriticCfg = field(default_factory=RslRlFlashSacCriticCfg)
"""The critic configuration."""
algorithm: RslRlFlashSacAlgorithmCfg = field(
default_factory=RslRlFlashSacAlgorithmCfg
)
"""The algorithm configuration."""
replay: RslRlReplayBufferCfg = field(default_factory=RslRlReplayBufferCfg)
"""The replay-buffer configuration."""
70 changes: 69 additions & 1 deletion src/mjlab/rl/runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@

import torch
from rsl_rl.env import VecEnv
from rsl_rl.runners import OnPolicyRunner
from rsl_rl.runners import OffPolicyRunner, OnPolicyRunner

from mjlab.rl.vecenv_wrapper import RslRlVecEnvWrapper

Expand Down Expand Up @@ -139,3 +139,71 @@ def load(
if infos and "env_state" in infos:
self.env.unwrapped.common_step_counter = infos["env_state"]["common_step_counter"]
return infos


class MjlabOffPolicyRunner(OffPolicyRunner):
"""Off-policy (FlashSAC) runner that persists environment state across checkpoints.

Parallels :class:`MjlabOnPolicyRunner` for the off-policy loop: it persists the
environment's ``common_step_counter``, gates W&B uploads on ``upload_model``,
and exports ONNX via the legacy (``dynamo=False``) path. FlashSAC checkpoints
are native to rsl-rl>=5, so no legacy key migration is needed.
"""

env: RslRlVecEnvWrapper

def export_policy_to_onnx(
self, path: str, filename: str = "policy.onnx", verbose: bool = False
) -> None:
"""Export policy to ONNX using the legacy export path (dynamo=False)."""
onnx_model = self.alg.get_policy().as_onnx(verbose=verbose)
onnx_model.to("cpu")
onnx_model.eval()
os.makedirs(path, exist_ok=True)
torch.onnx.export(
onnx_model,
onnx_model.get_dummy_inputs(), # type: ignore[operator]
os.path.join(path, filename),
export_params=True,
opset_version=18,
verbose=verbose,
input_names=onnx_model.input_names, # type: ignore[arg-type]
output_names=onnx_model.output_names, # type: ignore[arg-type]
dynamic_axes={},
dynamo=False,
)

@staticmethod
def _get_export_paths(checkpoint_path: str) -> tuple[Path, str, Path]:
"""Resolve ONNX export paths from a checkpoint path."""
export_dir = Path(checkpoint_path).parent
filename = f"{export_dir.name}.onnx"
return export_dir, filename, export_dir / filename

def save(self, path: str, infos=None) -> None:
"""Save checkpoint, persisting common_step_counter and gating W&B upload."""
env_state = {"common_step_counter": self.env.unwrapped.common_step_counter}
infos = {**(infos or {}), "env_state": env_state}
saved_dict = self.alg.save()
saved_dict["iter"] = self.current_learning_iteration
saved_dict["infos"] = infos
torch.save(saved_dict, path)
if self.cfg["upload_model"]:
self.logger.save_model(path, self.current_learning_iteration)

def load(
self,
path: str,
load_cfg: dict | None = None,
strict: bool = True,
map_location: str | None = None,
) -> dict:
"""Load checkpoint and restore common_step_counter."""
loaded_dict = torch.load(path, map_location=map_location, weights_only=False)
load_iteration = self.alg.load(loaded_dict, load_cfg, strict)
if load_iteration:
self.current_learning_iteration = loaded_dict["iter"]
infos = loaded_dict["infos"]
if infos and "env_state" in infos:
self.env.unwrapped.common_step_counter = infos["env_state"]["common_step_counter"]
return infos
21 changes: 19 additions & 2 deletions src/mjlab/tasks/velocity/config/g1/__init__.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,11 @@
from mjlab.tasks.registry import register_mjlab_task
from mjlab.tasks.velocity.rl import VelocityOnPolicyRunner
from mjlab.tasks.velocity.rl import VelocityOffPolicyRunner, VelocityOnPolicyRunner

from .env_cfgs import (
unitree_g1_flat_env_cfg,
unitree_g1_rough_env_cfg,
)
from .rl_cfg import unitree_g1_ppo_runner_cfg
from .rl_cfg import unitree_g1_flashsac_runner_cfg, unitree_g1_ppo_runner_cfg

register_mjlab_task(
task_id="Mjlab-Velocity-Rough-Unitree-G1",
Expand All @@ -22,3 +22,20 @@
rl_cfg=unitree_g1_ppo_runner_cfg(),
runner_cls=VelocityOnPolicyRunner,
)


def _unitree_g1_flat_flashsac_env_cfg(play: bool = False):
"""Flat G1 velocity env cfg with the default 512-env count for FlashSAC."""
cfg = unitree_g1_flat_env_cfg(play=play)
if not play:
cfg.scene.num_envs = 512
return cfg


register_mjlab_task(
task_id="Mjlab-Velocity-Flat-Unitree-G1-FlashSAC",
env_cfg=_unitree_g1_flat_flashsac_env_cfg(),
play_env_cfg=_unitree_g1_flat_flashsac_env_cfg(play=True),
rl_cfg=unitree_g1_flashsac_runner_cfg(),
runner_cls=VelocityOffPolicyRunner,
)
54 changes: 54 additions & 0 deletions src/mjlab/tasks/velocity/config/g1/rl_cfg.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,14 @@
"""RL configuration for Unitree G1 velocity task."""

from mjlab.rl import (
RslRlFlashSacActorCfg,
RslRlFlashSacAlgorithmCfg,
RslRlFlashSacCriticCfg,
RslRlModelCfg,
RslRlOffPolicyRunnerCfg,
RslRlOnPolicyRunnerCfg,
RslRlPpoAlgorithmCfg,
RslRlReplayBufferCfg,
)


Expand Down Expand Up @@ -44,3 +49,52 @@ def unitree_g1_ppo_runner_cfg() -> RslRlOnPolicyRunnerCfg:
num_steps_per_env=24,
max_iterations=30_000,
)


def unitree_g1_flashsac_runner_cfg() -> RslRlOffPolicyRunnerCfg:
"""Create the FlashSAC (off-policy) runner configuration for the G1 velocity task."""
return RslRlOffPolicyRunnerCfg(
actor=RslRlFlashSacActorCfg(
num_blocks=2,
# FlashSAC paper Table 9 (GPU sims): actor hidden 128, critic hidden 256.
hidden_dim=128,
obs_normalization=False,
),
critic=RslRlFlashSacCriticCfg(
num_blocks=2,
hidden_dim=256,
num_bins=101,
min_v=-5.0,
max_v=5.0,
num_qs=2,
obs_normalization=False,
),
algorithm=RslRlFlashSacAlgorithmCfg(
gamma=0.99,
# FlashSAC's IsaacLab locomotion recipe uses 3-step returns.
n_step=3,
critic_target_update_tau=0.01,
num_bins=101,
min_v=-5.0,
max_v=5.0,
actor_update_period=2,
normalize_reward=True,
normalized_g_max=5.0,
# Schedule length is expressed in gradient steps
# (num_steps_per_env * updates_per_step * max_iterations).
learning_rate_decay_steps=60_000,
),
replay=RslRlReplayBufferCfg(
capacity=1_000_000,
min_length=10_000,
sample_batch_size=2048,
),
# Off-policy collection/update cadence. With ~512 envs this collects 512
# transitions per iteration and takes `num_steps_per_env * updates_per_step`
# gradient steps.
num_steps_per_env=1,
updates_per_step=2.0,
experiment_name="g1_velocity_flashsac",
save_interval=50,
max_iterations=30_000,
)
3 changes: 3 additions & 0 deletions src/mjlab/tasks/velocity/rl/__init__.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,6 @@
from mjlab.tasks.velocity.rl.runner import (
VelocityOffPolicyRunner as VelocityOffPolicyRunner,
)
from mjlab.tasks.velocity.rl.runner import (
VelocityOnPolicyRunner as VelocityOnPolicyRunner,
)
28 changes: 27 additions & 1 deletion src/mjlab/tasks/velocity/rl/runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
attach_metadata_to_onnx,
get_base_metadata,
)
from mjlab.rl.runner import MjlabOnPolicyRunner
from mjlab.rl.runner import MjlabOffPolicyRunner, MjlabOnPolicyRunner


class VelocityOnPolicyRunner(MjlabOnPolicyRunner):
Expand All @@ -30,3 +30,29 @@ def save(self, path: str, infos=None):
wandb.save(str(onnx_path), base_path=str(policy_dir))
except Exception as e:
print(f"[WARN] ONNX export failed (training continues): {e}")


class VelocityOffPolicyRunner(MjlabOffPolicyRunner):
"""FlashSAC velocity runner that also exports ONNX (+ metadata) on save."""

env: RslRlVecEnvWrapper

def save(self, path: str, infos=None):
super().save(path, infos)
policy_dir, filename, onnx_path = self._get_export_paths(path)
try:
self.export_policy_to_onnx(str(policy_dir), filename)
run_name: str = (
wandb.run.name
if self.logger.logger_type in ("wandb", "WandbLogWriter") and wandb.run
else "local"
) # type: ignore[assignment]
metadata = get_base_metadata(self.env.unwrapped, run_name)
attach_metadata_to_onnx(str(onnx_path), metadata)
if (
self.logger.logger_type in ("wandb", "WandbLogWriter")
and self.cfg["upload_model"]
):
wandb.save(str(onnx_path), base_path=str(policy_dir))
except Exception as e:
print(f"[WARN] ONNX export failed (training continues): {e}")