diff --git a/CITATION.cff b/CITATION.cff new file mode 100644 index 0000000..8c21518 --- /dev/null +++ b/CITATION.cff @@ -0,0 +1,29 @@ +cff-version: 1.2.0 +message: "If you use this code, please cite the paper below." +preferred-citation: + type: article + title: "Policy-Driven World Model Adaptation for Robust Offline Model-based Reinforcement Learning" + authors: + - family-names: Chen + given-names: Jiayu + - family-names: Xu + given-names: Le + - family-names: Venugopal + given-names: Aravind + - family-names: Schneider + given-names: Jeff + year: 2025 + journal: "arXiv preprint arXiv:2505.13709" + url: "https://arxiv.org/abs/2505.13709" +title: "ROMBRL: Policy-Driven World Model Adaptation for Robust Offline Model-based Reinforcement Learning" +authors: + - family-names: Chen + given-names: Jiayu + - family-names: Xu + given-names: Le + - family-names: Venugopal + given-names: Aravind + - family-names: Schneider + given-names: Jeff +url: "https://github.com/Agentic-Intelligence-Lab/ROMBRL" +license: MIT diff --git a/D4RL/.dockerignore b/D4RL/.dockerignore new file mode 100644 index 0000000..e053ae7 --- /dev/null +++ b/D4RL/.dockerignore @@ -0,0 +1,28 @@ +__pycache__/ +*.py[cod] +.venv/ +venv/ +*.egg-info/ +log/ +logs/ +runs/ +wandb/ +tensorboard/ +*.log +data/ +datasets/ +checkpoints/ +checkpoint/ +models/ +*.h5 +*.hdf5 +*.npy +*.npz +*.pkl +*.pt +*.pth +*.ckpt +build/ +dist/ +*.so +.git/ diff --git a/D4RL/Dockerfile b/D4RL/Dockerfile new file mode 100644 index 0000000..544f1f3 --- /dev/null +++ b/D4RL/Dockerfile @@ -0,0 +1,47 @@ +# ROMBRL — D4RL MuJoCo experiments (Tables 1, 2, 4) +# +# NOTE: this Dockerfile has been written to match D4RL/requirements.txt and +# D4RL/README.md but has not been build-tested in this environment (no local +# Docker daemon available). Build and smoke-test before relying on it: +# docker build -t rombrl-d4rl -f D4RL/Dockerfile D4RL/ # context must be D4RL/, not repo root +# docker run --rm -it rombrl-d4rl python run_rombrl2.py --help + +FROM pytorch/pytorch:2.0.1-cuda11.7-cudnn8-devel + +ENV DEBIAN_FRONTEND=noninteractive \ + MUJOCO_PY_MUJOCO_PATH=/root/.mujoco/mujoco210 \ + LD_LIBRARY_PATH=/root/.mujoco/mujoco210/bin:${LD_LIBRARY_PATH} + +# System dependencies for mujoco-py / gym mujoco rendering + Cython build tools +RUN apt-get update && apt-get install -y --no-install-recommends \ + build-essential \ + git \ + wget \ + unzip \ + patchelf \ + libosmesa6-dev \ + libgl1-mesa-glx \ + libgl1-mesa-dev \ + libglfw3 \ + libglew-dev \ + libglu1-mesa \ + libglu1-mesa-dev \ + && rm -rf /var/lib/apt/lists/* + +# MuJoCo 2.1.0 binary (required by mujoco-py) +RUN mkdir -p /root/.mujoco && \ + wget -q https://github.com/google-deepmind/mujoco/releases/download/2.1.0/mujoco210-linux-x86_64.tar.gz -O /tmp/mujoco210.tar.gz && \ + tar -xzf /tmp/mujoco210.tar.gz -C /root/.mujoco && \ + rm /tmp/mujoco210.tar.gz + +WORKDIR /workspace/D4RL + +COPY requirements.txt . +RUN pip install --no-cache-dir -r requirements.txt + +COPY . . + +# Build the ctree Cython/C++ extension used by the search-based baselines +RUN cd offlinerlkit/utils/ctree && python setup.py build_ext --inplace && cd ../../.. + +CMD ["/bin/bash"] diff --git a/D4RL/README.md b/D4RL/README.md index 9c11a15..8697f98 100644 --- a/D4RL/README.md +++ b/D4RL/README.md @@ -1,8 +1,17 @@ # Policy-Driven World Model Adaptation for Robust Offline Model-based Reinforcement Learning -- Please download the dynamics/reward models and hyperparameter files from [d4rl_data](https://drive.google.com/drive/folders/1FiJbpAJvul629u4VjgOyugHwcBPJyc7u?usp=sharing) to the folder `data`. +## Setup -- Please set up a virtual environment based on the instructions from [OfflineRLKit](https://github.com/yihaosun1124/OfflineRL-Kit). +- Download the mujoco210 binary from the [MuJoCo 2.1.0 release](https://github.com/google-deepmind/mujoco/releases/tag/2.1.0) and place it at `~/.mujoco/mujoco210` (required by `mujoco-py`/`gym.envs.mujoco`). + +- Create a Python virtual environment (Python 3.8–3.9 recommended) and install dependencies: + ```bash + pip install -r requirements.txt + cd offlinerlkit/utils/ctree && python setup.py build_ext --inplace && cd ../../.. + ``` + This codebase builds on [OfflineRL-Kit](https://github.com/yihaosun1124/OfflineRL-Kit); consult it if you hit environment issues not covered here. + +- Download the dynamics/reward models and hyperparameter files from [d4rl_data](https://drive.google.com/drive/folders/1FiJbpAJvul629u4VjgOyugHwcBPJyc7u?usp=sharing) to the folder `data`. ## Table 1: Noisy D4RL MuJoCo @@ -18,7 +27,7 @@ python run_XXX.py rombrl2, cql, edac, combo, rambo, mobile, rorl, tracer, rfqi ``` -These correspond to ROMBRL and the baselines reported in Table 1. The `run_bambrl.py` script is also kept in the repository for additional/legacy comparisons. +These correspond to ROMBRL and the baselines reported in Table 1. The `run_bamcts.py` script is also kept in the repository for additional/legacy comparisons. To specify the task and random seed for each run, change `load_path_id` at the bottom of each `run_XXX.py` script. The default D4RL evaluation uses measurement noise controlled by `--noise_scale`; the Table 1 noisy setting uses `--noise_scale 0.05`. diff --git a/D4RL/offlinerlkit/policy/__init__.py b/D4RL/offlinerlkit/policy/__init__.py index b2e4288..2700f11 100644 --- a/D4RL/offlinerlkit/policy/__init__.py +++ b/D4RL/offlinerlkit/policy/__init__.py @@ -16,7 +16,7 @@ from offlinerlkit.policy.model_based.mobile import MOBILEPolicy from offlinerlkit.policy.model_based.rambo import RAMBOPolicy from offlinerlkit.policy.model_based.combo import COMBOPolicy -from offlinerlkit.policy.model_based.bambrl import BAMBRLPolicy +from offlinerlkit.policy.model_based.bamcts import BAMCTSPolicy #from offlinerlkit.policy.model_based.rombrl import ROMBRLPolicy from offlinerlkit.policy.model_based.rombrl2 import ROMBRL2Policy #from offlinerlkit.policy.model_based.rombrl3 import ROMBRL3Policy @@ -37,7 +37,7 @@ "MOBILEPolicy", "RAMBOPolicy", "COMBOPolicy", - "BAMBRLPolicy", + "BAMCTSPolicy", #"ROMBRLPolicy", "ROMBRL2Policy", #"ROMBRL3Policy" diff --git a/D4RL/offlinerlkit/policy/model_based/bambrl.py b/D4RL/offlinerlkit/policy/model_based/bamcts.py similarity index 99% rename from D4RL/offlinerlkit/policy/model_based/bambrl.py rename to D4RL/offlinerlkit/policy/model_based/bamcts.py index 2c5aed0..59ec575 100644 --- a/D4RL/offlinerlkit/policy/model_based/bambrl.py +++ b/D4RL/offlinerlkit/policy/model_based/bamcts.py @@ -11,7 +11,7 @@ from offlinerlkit.buffer import SLReplayBuffer, SL_Transition from torch.distributions import Normal, Independent -class BAMBRLPolicy(MOBILEPolicy): +class BAMCTSPolicy(MOBILEPolicy): def __init__( self, diff --git a/D4RL/offlinerlkit/policy/model_free/rorl_old.py b/D4RL/offlinerlkit/policy/model_free/rorl_old.py deleted file mode 100644 index abd6f31..0000000 --- a/D4RL/offlinerlkit/policy/model_free/rorl_old.py +++ /dev/null @@ -1,255 +0,0 @@ -import numpy as np -import torch -import torch.nn as nn -from torch.distributions import kl_divergence -from copy import deepcopy -from typing import Dict, Union, Tuple -from offlinerlkit.policy import BasePolicy -from offlinerlkit.utils.scaler import StandardScaler - -class RORLPolicy(BasePolicy): - """ - Robust Offline Reinforcement Learning (RORL) via Conservative Smoothing - """ - - def __init__( - self, - actor: nn.Module, - critics: nn.ModuleList, - actor_optim: torch.optim.Optimizer, - critics_optim: torch.optim.Optimizer, - tau: float = 0.005, - gamma: float = 0.99, - alpha: Union[float, Tuple[float, torch.Tensor, torch.optim.Optimizer]] = 0.2, - max_q_backup: bool = False, - deterministic_backup: bool = False, # RORL Default is False - # RORL Specific Hyperparameters - num_samples: int = 20, - policy_smooth_eps: float = 0.0, - policy_smooth_reg: float = 0.0, - q_smooth_eps: float = 0.0, - q_smooth_reg: float = 0.0, - q_smooth_tau: float = 0.2, - obs_std: float = 1.0, # Used for scaling noise - scaler: StandardScaler = None, # Input Normalizer - device: str = "cpu" - ) -> None: - - super().__init__() - self.actor = actor - self.critics = critics - self.critics_old = deepcopy(critics) - self.critics_old.eval() - - self.actor_optim = actor_optim - self.critics_optim = critics_optim - - self._tau = tau - self._gamma = gamma - self.device = device - - self._is_auto_alpha = False - if isinstance(alpha, tuple): - self._is_auto_alpha = True - self._target_entropy, self._log_alpha, self.alpha_optim = alpha - self._alpha = self._log_alpha.detach().exp() - else: - self._alpha = alpha - - self._max_q_backup = max_q_backup - self._deterministic_backup = deterministic_backup - - # RORL Parameters - self.num_samples = num_samples - self.policy_smooth_eps = policy_smooth_eps - self.policy_smooth_reg = policy_smooth_reg - self.q_smooth_eps = q_smooth_eps - self.q_smooth_reg = q_smooth_reg - self.q_smooth_tau = q_smooth_tau - self.obs_std = obs_std - - # Normalization - self.scaler = scaler - - def train(self) -> None: - self.actor.train() - self.critics.train() - - def eval(self) -> None: - self.actor.eval() - self.critics.eval() - - def _sync_weight(self) -> None: - for o, n in zip(self.critics_old.parameters(), self.critics.parameters()): - o.data.copy_(o.data * (1.0 - self._tau) + n.data * self._tau) - - def actforward( - self, - obs: torch.Tensor, - deterministic: bool = False - ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - # NOTE: This function expects `obs` to be ALREADY normalized if scaler exists. - dist = self.actor(obs) - if deterministic: - squashed_action, raw_action = dist.mode() - else: - squashed_action, raw_action = dist.rsample() - log_prob = dist.log_prob(squashed_action, raw_action) - return squashed_action, log_prob, dist.mean, dist.stddev - - def select_action( - self, - obs: np.ndarray, - deterministic: bool = False - ) -> np.ndarray: - # === Evaluation Logic === - # 1. Input `obs` is RAW state (potentially with physics noise from Robust Eval) - # 2. Apply Normalization (if enabled) - if self.scaler is not None: - obs = self.scaler.transform(obs) - - # 3. To Tensor & Inference - with torch.no_grad(): - obs = torch.FloatTensor(obs).to(self.device).unsqueeze(0) - action, _, _, _ = self.actforward(obs, deterministic) - - return action.cpu().numpy()[0] - - def _get_noised_obs(self, obs, eps): - # Ported from RORL: trainers/q_learning/sac.py - # Noise is Uniform[-eps*std, eps*std] - M, N = obs.shape[0], obs.shape[1] - size = self.num_samples - - delta_s = 2 * eps * self.obs_std * (torch.rand(size, N, device=self.device) - 0.5) - - # Expand obs: (M, N) -> (M*size, N) - tmp_obs = obs.reshape(-1, 1, N).repeat(1, size, 1).reshape(-1, N) - - # Expand noise: (size, N) -> (M*size, N) - delta_s = delta_s.reshape(1, size, N).repeat(M, 1, 1).reshape(-1, N) - - noised_obs = tmp_obs + delta_s - return M, size, noised_obs - - def learn(self, batch: Dict) -> Dict: - obss, actions, next_obss, rewards, terminals = \ - batch["observations"], batch["actions"], batch["next_observations"], batch["rewards"], batch["terminals"] - - # === Training Normalization Logic === - # Ensure we train on normalized data if scaler is present - if self.scaler is not None: - obss = self.scaler.transform_tensor(obss) - next_obss = self.scaler.transform_tensor(next_obss) - - batch_size = obss.shape[0] - action_dim = actions.shape[-1] - - # ------------------------- - # 1. Update Actor - # ------------------------- - a, log_probs, policy_mean, policy_std = self.actforward(obss) - qas = self.critics(obss, a) - actor_loss = -torch.min(qas, 0)[0].mean() + self._alpha * log_probs.mean() - - # === RORL: Policy Smoothing === - if self.policy_smooth_eps > 0 and self.policy_smooth_reg > 0: - M, size, noised_obs = self._get_noised_obs(obss, self.policy_smooth_eps) - - # Get policy distribution on noisy states - _, _, noised_policy_mean, noised_policy_std = self.actforward(noised_obs) - - # Construct distributions for KL calculation - # We need to expand original mean/std to match the noisy samples size (M*size) - orig_mean_expanded = policy_mean.reshape(-1, 1, action_dim).repeat(1, size, 1).reshape(-1, action_dim) - orig_std_expanded = policy_std.reshape(-1, 1, action_dim).repeat(1, size, 1).reshape(-1, action_dim) - - action_dist = torch.distributions.Normal(orig_mean_expanded, orig_std_expanded) - noised_action_dist = torch.distributions.Normal(noised_policy_mean, noised_policy_std) - - # Symmetric KL Divergence - kl_loss = kl_divergence(action_dist, noised_action_dist).sum(dim=-1) + \ - kl_divergence(noised_action_dist, action_dist).sum(dim=-1) - - kl_loss = kl_loss.reshape(M, size) - # Take max KL over the N samples for each transition - kl_loss_max = kl_loss.max(dim=1)[0].mean() - - actor_loss += self.policy_smooth_reg * kl_loss_max - - self.actor_optim.zero_grad() - actor_loss.backward() - self.actor_optim.step() - - if self._is_auto_alpha: - log_probs = log_probs.detach() + self._target_entropy - alpha_loss = -(self._log_alpha * log_probs).mean() - self.alpha_optim.zero_grad() - alpha_loss.backward() - self.alpha_optim.step() - self._alpha = torch.clamp(self._log_alpha.detach().exp(), 0.0, 1.0) - - # ------------------------- - # 2. Update Critic - # ------------------------- - if self._max_q_backup: - with torch.no_grad(): - tmp_next_obss = next_obss.unsqueeze(1).repeat(1, 10, 1) \ - .view(batch_size * 10, next_obss.shape[-1]) - tmp_next_actions, _, _, _ = self.actforward(tmp_next_obss) - tmp_next_qs = self.critics_old(tmp_next_obss, tmp_next_actions) \ - .view(self.critics._num_ensemble, batch_size, 10, 1).max(2)[0] \ - .view(self.critics._num_ensemble, batch_size, 1) - next_q = tmp_next_qs.min(0)[0] - else: - with torch.no_grad(): - next_actions, next_log_probs, _, _ = self.actforward(next_obss) - next_q = self.critics_old(next_obss, next_actions).min(0)[0] - if not self._deterministic_backup: - next_q -= self._alpha * next_log_probs - - target_q = rewards + self._gamma * (1 - terminals) * next_q - qs = self.critics(obss, actions) - critics_loss = ((qs - target_q.unsqueeze(0)).pow(2)).mean(dim=(1, 2)).sum() - - # === RORL: Q-function Smoothing === - if self.q_smooth_eps > 0 and self.q_smooth_reg > 0: - M, size, noised_obs = self._get_noised_obs(obss, self.q_smooth_eps) - - # Re-use actions for the noisy states: (M, A) -> (M*size, A) - actions_repeated = actions.reshape(-1, 1, action_dim).repeat(1, size, 1).reshape(-1, action_dim) - - noised_qs_pred = self.critics(noised_obs, actions_repeated) # [num_critics, M*size, 1] - - # Expand original Qs to compare: [num_critics, M, 1] -> [num_critics, M*size, 1] - qs_pred_expanded = qs.repeat(1, 1, size).reshape(self.critics._num_ensemble, -1, 1) - - diff = noised_qs_pred - qs_pred_expanded - - # Asymmetric Loss - zero_tensor = torch.zeros_like(diff) - pos = torch.maximum(diff, zero_tensor) - neg = torch.minimum(diff, zero_tensor) - - noise_Q_loss = (1 - self.q_smooth_tau) * pos.pow(2) + self.q_smooth_tau * neg.pow(2) - noise_Q_loss = noise_Q_loss.mean(dim=0).reshape(M, size) - noise_Q_loss_max = noise_Q_loss.max(dim=1)[0].mean() - - critics_loss += self.q_smooth_reg * noise_Q_loss_max - - self.critics_optim.zero_grad() - critics_loss.backward() - self.critics_optim.step() - - self._sync_weight() - - result = { - "loss/actor": actor_loss.item(), - "loss/critics": critics_loss.item() - } - - if self._is_auto_alpha: - result["loss/alpha"] = alpha_loss.item() - result["alpha"] = self._alpha.item() - - return result \ No newline at end of file diff --git a/D4RL/offlinerlkit/policy/model_free/rorl_old2.py b/D4RL/offlinerlkit/policy/model_free/rorl_old2.py deleted file mode 100644 index d0342f2..0000000 --- a/D4RL/offlinerlkit/policy/model_free/rorl_old2.py +++ /dev/null @@ -1,268 +0,0 @@ -import numpy as np -import torch -import torch.nn as nn -from torch.distributions import kl_divergence -from copy import deepcopy -from typing import Dict, Union, Tuple -from offlinerlkit.policy import BasePolicy -from offlinerlkit.utils.scaler import StandardScaler - -class RORLPolicy(BasePolicy): - """ - Robust Offline Reinforcement Learning (RORL) via Conservative Smoothing - """ - - def __init__( - self, - actor: nn.Module, - critics: nn.ModuleList, - actor_optim: torch.optim.Optimizer, - critics_optim: torch.optim.Optimizer, - tau: float = 0.005, - gamma: float = 0.99, - alpha: Union[float, Tuple[float, torch.Tensor, torch.optim.Optimizer]] = 0.2, - max_q_backup: bool = False, - deterministic_backup: bool = False, - # === RORL: Consistency / Smoothing Params === - num_samples: int = 20, - policy_smooth_eps: float = 0.0, - policy_smooth_reg: float = 0.0, - q_smooth_eps: float = 0.0, - q_smooth_reg: float = 0.0, - q_smooth_tau: float = 0.2, - # === RORL: OOD Conservative Params (New) === - q_ood_eps: float = 0.0, - q_ood_reg: float = 0.0, - q_ood_uncertainty_reg: float = 0.0, - q_ood_uncertainty_reg_min: float = 0.0, - # =========================================== - obs_std: float = 1.0, - scaler: StandardScaler = None, - device: str = "cpu" - ) -> None: - - super().__init__() - self.actor = actor - self.critics = critics - self.critics_old = deepcopy(critics) - self.critics_old.eval() - - self.actor_optim = actor_optim - self.critics_optim = critics_optim - - self._tau = tau - self._gamma = gamma - self.device = device - - self._is_auto_alpha = False - if isinstance(alpha, tuple): - self._is_auto_alpha = True - self._target_entropy, self._log_alpha, self.alpha_optim = alpha - self._alpha = self._log_alpha.detach().exp() - else: - self._alpha = alpha - - self._max_q_backup = max_q_backup - self._deterministic_backup = deterministic_backup - - # RORL Parameters - self.num_samples = num_samples - self.policy_smooth_eps = policy_smooth_eps - self.policy_smooth_reg = policy_smooth_reg - self.q_smooth_eps = q_smooth_eps - self.q_smooth_reg = q_smooth_reg - self.q_smooth_tau = q_smooth_tau - - # OOD Parameters - self.q_ood_eps = q_ood_eps - self.q_ood_reg = q_ood_reg - self.q_ood_uncertainty_reg = q_ood_uncertainty_reg - self.q_ood_uncertainty_reg_min = q_ood_uncertainty_reg_min - - self.obs_std = obs_std - self.scaler = scaler - - def train(self) -> None: - self.actor.train() - self.critics.train() - - def eval(self) -> None: - self.actor.eval() - self.critics.eval() - - def _sync_weight(self) -> None: - for o, n in zip(self.critics_old.parameters(), self.critics.parameters()): - o.data.copy_(o.data * (1.0 - self._tau) + n.data * self._tau) - - def actforward( - self, - obs: torch.Tensor, - deterministic: bool = False - ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - dist = self.actor(obs) - if deterministic: - squashed_action, raw_action = dist.mode() - else: - squashed_action, raw_action = dist.rsample() - log_prob = dist.log_prob(squashed_action, raw_action) - return squashed_action, log_prob, dist.mean, dist.stddev - - def select_action( - self, - obs: np.ndarray, - deterministic: bool = False - ) -> np.ndarray: - if self.scaler is not None: - obs = self.scaler.transform(obs) - - with torch.no_grad(): - obs = torch.FloatTensor(obs).to(self.device).unsqueeze(0) - action, _, _, _ = self.actforward(obs, deterministic) - - return action.cpu().numpy()[0] - - def _get_noised_obs(self, obs, eps): - # Noise is Uniform[-eps*std, eps*std] - M, N = obs.shape[0], obs.shape[1] - size = self.num_samples - - delta_s = 2 * eps * self.obs_std * (torch.rand(size, N, device=self.device) - 0.5) - - tmp_obs = obs.reshape(-1, 1, N).repeat(1, size, 1).reshape(-1, N) - delta_s = delta_s.reshape(1, size, N).repeat(M, 1, 1).reshape(-1, N) - - noised_obs = tmp_obs + delta_s - return M, size, noised_obs - - def learn(self, batch: Dict) -> Dict: - obss, actions, next_obss, rewards, terminals = \ - batch["observations"], batch["actions"], batch["next_observations"], batch["rewards"], batch["terminals"] - - if self.scaler is not None: - obss = self.scaler.transform_tensor(obss) - next_obss = self.scaler.transform_tensor(next_obss) - - batch_size = obss.shape[0] - action_dim = actions.shape[-1] - - # ------------------------- - # 1. Update Actor - # ------------------------- - a, log_probs, policy_mean, policy_std = self.actforward(obss) - qas = self.critics(obss, a) - actor_loss = -torch.min(qas, 0)[0].mean() + self._alpha * log_probs.mean() - - # === RORL: Policy Smoothing === - if self.policy_smooth_eps > 0 and self.policy_smooth_reg > 0: - M, size, noised_obs = self._get_noised_obs(obss, self.policy_smooth_eps) - _, _, noised_policy_mean, noised_policy_std = self.actforward(noised_obs) - - orig_mean_expanded = policy_mean.reshape(-1, 1, action_dim).repeat(1, size, 1).reshape(-1, action_dim) - orig_std_expanded = policy_std.reshape(-1, 1, action_dim).repeat(1, size, 1).reshape(-1, action_dim) - - action_dist = torch.distributions.Normal(orig_mean_expanded, orig_std_expanded) - noised_action_dist = torch.distributions.Normal(noised_policy_mean, noised_policy_std) - - kl_loss = kl_divergence(action_dist, noised_action_dist).sum(dim=-1) + \ - kl_divergence(noised_action_dist, action_dist).sum(dim=-1) - kl_loss = kl_loss.reshape(M, size) - kl_loss_max = kl_loss.max(dim=1)[0].mean() - - actor_loss += self.policy_smooth_reg * kl_loss_max - - self.actor_optim.zero_grad() - actor_loss.backward() - self.actor_optim.step() - - if self._is_auto_alpha: - log_probs = log_probs.detach() + self._target_entropy - alpha_loss = -(self._log_alpha * log_probs).mean() - self.alpha_optim.zero_grad() - alpha_loss.backward() - self.alpha_optim.step() - self._alpha = torch.clamp(self._log_alpha.detach().exp(), 0.0, 1.0) - - # ------------------------- - # 2. Update Critic - # ------------------------- - if self._max_q_backup: - with torch.no_grad(): - tmp_next_obss = next_obss.unsqueeze(1).repeat(1, 10, 1) \ - .view(batch_size * 10, next_obss.shape[-1]) - tmp_next_actions, _, _, _ = self.actforward(tmp_next_obss) - tmp_next_qs = self.critics_old(tmp_next_obss, tmp_next_actions) \ - .view(self.critics._num_ensemble, batch_size, 10, 1).max(2)[0] \ - .view(self.critics._num_ensemble, batch_size, 1) - next_q = tmp_next_qs.min(0)[0] - else: - with torch.no_grad(): - next_actions, next_log_probs, _, _ = self.actforward(next_obss) - next_q = self.critics_old(next_obss, next_actions).min(0)[0] - if not self._deterministic_backup: - next_q -= self._alpha * next_log_probs - - target_q = rewards + self._gamma * (1 - terminals) * next_q - qs = self.critics(obss, actions) - critics_loss = ((qs - target_q.unsqueeze(0)).pow(2)).mean(dim=(1, 2)).sum() - - # === RORL: Q-function Smoothing (Consistency) === - if self.q_smooth_eps > 0 and self.q_smooth_reg > 0: - M, size, noised_obs = self._get_noised_obs(obss, self.q_smooth_eps) - actions_repeated = actions.reshape(-1, 1, action_dim).repeat(1, size, 1).reshape(-1, action_dim) - noised_qs_pred = self.critics(noised_obs, actions_repeated) - qs_pred_expanded = qs.repeat(1, 1, size).reshape(self.critics._num_ensemble, -1, 1) - - diff = noised_qs_pred - qs_pred_expanded - pos = torch.maximum(diff, torch.zeros_like(diff)) - neg = torch.minimum(diff, torch.zeros_like(diff)) - - noise_Q_loss = (1 - self.q_smooth_tau) * pos.pow(2) + self.q_smooth_tau * neg.pow(2) - noise_Q_loss = noise_Q_loss.mean(dim=0).reshape(M, size) - noise_Q_loss_max = noise_Q_loss.max(dim=1)[0].mean() - - critics_loss += self.q_smooth_reg * noise_Q_loss_max - - # === RORL: OOD Conservative Penalty (Value Suppression & Uncertainty) === - # Corresponds to q_ood_reg and q_ood_uncertainty_reg in sac.py - if self.q_ood_eps > 0 and (self.q_ood_reg > 0 or self.q_ood_uncertainty_reg > 0): - # 1. Sample OOD states (typically with larger epsilon than smoothing, or same) - M, size, ood_obs = self._get_noised_obs(obss, self.q_ood_eps) - - # 2. Sample actions from the current policy on these OOD states - ood_actions, _, _, _ = self.actforward(ood_obs) - - # 3. Calculate Q values for OOD (s', a') - ood_qs = self.critics(ood_obs, ood_actions) # Shape: [num_critics, M*size, 1] - - # 4. Conservative Value Penalty: Minimize mean Q value on OOD states - if self.q_ood_reg > 0: - # We typically want to minimize the Q values - critics_loss += self.q_ood_reg * ood_qs.mean() - - # 5. Uncertainty Penalty: Minimize variance/std of Q values on OOD states - if self.q_ood_uncertainty_reg > 0: - # Calculate std across the ensemble dimension (dim=0) - ood_qs_std = ood_qs.std(dim=0) # Shape: [M*size, 1] - - # Apply optional hinge loss for uncertainty (std - min)+ - if self.q_ood_uncertainty_reg_min > 0: - ood_qs_std = torch.clamp(ood_qs_std - self.q_ood_uncertainty_reg_min, min=0) - - critics_loss += self.q_ood_uncertainty_reg * ood_qs_std.mean() - - self.critics_optim.zero_grad() - critics_loss.backward() - self.critics_optim.step() - - self._sync_weight() - - result = { - "loss/actor": actor_loss.item(), - "loss/critics": critics_loss.item() - } - - if self._is_auto_alpha: - result["loss/alpha"] = alpha_loss.item() - result["alpha"] = self._alpha.item() - - return result \ No newline at end of file diff --git a/D4RL/requirements.txt b/D4RL/requirements.txt new file mode 100644 index 0000000..ed267d1 --- /dev/null +++ b/D4RL/requirements.txt @@ -0,0 +1,27 @@ +# Dependencies for the D4RL MuJoCo experiments (Tables 1, 2, 4). +# Versions are inferred from the code's imports and known-compatible combinations +# for this MuJoCo/D4RL/gym vintage; if you hit a version conflict in your own +# environment, `pip freeze` after a successful install and pin exact versions here. +# +# MuJoCo itself is NOT installable via pip: download the mujoco210 binary from +# https://github.com/google-deepmind/mujoco/releases/tag/2.1.0 and place it at +# ~/.mujoco/mujoco210 (required by mujoco-py / gym.envs.mujoco) before installing +# the packages below. + +torch>=1.13,<2.1 +numpy>=1.21,<1.24 +scipy>=1.7 +pandas>=1.3 +matplotlib>=3.5 +easydict +tqdm +Cython>=0.29,<3.0 + +# RL environments +gym==0.23.1 +gymnasium>=0.28 +mujoco-py>=2.1,<2.2 +d4rl @ git+https://github.com/Farama-Foundation/d4rl@master#egg=d4rl + +# Baseline algorithm dependencies +stable-baselines3>=1.6,<2.0 diff --git a/D4RL/run_bambrl.py b/D4RL/run_bamcts.py similarity index 98% rename from D4RL/run_bambrl.py rename to D4RL/run_bamcts.py index 2d36bc3..870d7b4 100644 --- a/D4RL/run_bambrl.py +++ b/D4RL/run_bamcts.py @@ -19,13 +19,13 @@ from offlinerlkit.buffer import ReplayBuffer, BayesReplayBuffer, SLReplayBuffer from offlinerlkit.utils.logger import Logger, make_log_dirs from offlinerlkit.policy_trainer import BayesMBPolicyTrainer -from offlinerlkit.policy import BAMBRLPolicy +from offlinerlkit.policy import BAMCTSPolicy from offlinerlkit.utils.searcher import Searcher def get_args(): parser = argparse.ArgumentParser() - parser.add_argument("--algo-name", type=str, default="bambrl") + parser.add_argument("--algo-name", type=str, default="bamcts") parser.add_argument("--task", type=str, default="walker2d-medium-expert-v2") parser.add_argument("--seed", type=int, default=1) parser.add_argument("--actor-lr", type=float, default=1e-4) @@ -273,7 +273,7 @@ def train(load_path=None, eval_path=None): searcher = None # create policy - policy = BAMBRLPolicy( + policy = BAMCTSPolicy( args.elite_only, elite_list, args.use_ba, @@ -374,5 +374,5 @@ def train(load_path=None, eval_path=None): '/data/wk-med/seed-0', '/data/wk-med/seed-1', '/data/wk-med/seed-2', '/data/wk-rnd/seed-0', '/data/wk-rnd/seed-1', '/data/wk-rnd/seed-2'] load_path_id = 25 # 0-6 - # eval_path = '/log/walker2d-medium-replay-v2/bambrl_mcts&penalty_coef=0.5&rollout_length=1&real_ratio=0.05/seed_1×tamp_24-1214-131101/checkpoint' + # eval_path = '/log/walker2d-medium-replay-v2/bamcts_mcts&penalty_coef=0.5&rollout_length=1&real_ratio=0.05/seed_1×tamp_24-1214-131101/checkpoint' train(current_working_directory + load_path_ls[load_path_id]) \ No newline at end of file diff --git a/Fusion/.dockerignore b/Fusion/.dockerignore new file mode 100644 index 0000000..e053ae7 --- /dev/null +++ b/Fusion/.dockerignore @@ -0,0 +1,28 @@ +__pycache__/ +*.py[cod] +.venv/ +venv/ +*.egg-info/ +log/ +logs/ +runs/ +wandb/ +tensorboard/ +*.log +data/ +datasets/ +checkpoints/ +checkpoint/ +models/ +*.h5 +*.hdf5 +*.npy +*.npz +*.pkl +*.pt +*.pth +*.ckpt +build/ +dist/ +*.so +.git/ diff --git a/Fusion/Dockerfile b/Fusion/Dockerfile new file mode 100644 index 0000000..034a625 --- /dev/null +++ b/Fusion/Dockerfile @@ -0,0 +1,38 @@ +# ROMBRL — Tokamak Control experiments (Table 3) +# +# NOTE: this Dockerfile has been written to match Fusion/requirements.txt and +# Fusion/README.md but has not been build-tested in this environment (no local +# Docker daemon available). Build and smoke-test before relying on it: +# docker build -t rombrl-fusion -f Fusion/Dockerfile Fusion/ # context must be Fusion/, not repo root +# docker run --rm -it rombrl-fusion python rl_scripts/run_rombrl.py --help +# +# DIII-D operational data is proprietary and is NOT included in this image — +# see Fusion/README.md. Mount your own data directory at runtime, e.g.: +# docker run --rm -it -v /path/to/your/data:/workspace/Fusion/data rombrl-fusion + +FROM pytorch/pytorch:2.0.1-cuda11.7-cudnn8-devel + +ENV DEBIAN_FRONTEND=noninteractive + +RUN apt-get update && apt-get install -y --no-install-recommends \ + build-essential \ + git \ + && rm -rf /var/lib/apt/lists/* + +WORKDIR /workspace/Fusion + +COPY requirements.txt . +RUN pip install --no-cache-dir -r requirements.txt + +# dynamics-toolbox is a separate repo required by the Fusion dynamics model +RUN git clone https://github.com/LucasCJYSDL/dynamics-toolbox.git /opt/dynamics-toolbox && \ + cd /opt/dynamics-toolbox && \ + pip install --no-cache-dir -r requirements.txt && \ + pip install --no-cache-dir -e . + +COPY . . + +# Build the ctree Cython/C++ extension used by the search-based baselines +RUN cd offlinerlkit/utils/ctree && python setup.py build_ext --inplace && cd ../../.. + +CMD ["/bin/bash"] diff --git a/Fusion/README.md b/Fusion/README.md index ada4590..61df6a5 100644 --- a/Fusion/README.md +++ b/Fusion/README.md @@ -2,12 +2,17 @@ ## Requirements -- You need to download/clone this open-source repo: [dynamics-toolbox](https://github.com/LucasCJYSDL/dynamics-toolbox). +- Create a Python 3.9 virtual environment and install this folder's dependencies: + ```bash + pip install -r requirements.txt + cd offlinerlkit/utils/ctree && python setup.py build_ext --inplace && cd ../../.. + ``` -- Make a virtual environment with python=3.9, then enter the repo above and run: +- Download/clone [dynamics-toolbox](https://github.com/LucasCJYSDL/dynamics-toolbox) and, in the **same** virtual environment, install it: ```bash + cd dynamics-toolbox pip install -r requirements.txt - pip install -e . + pip install -e . ``` ## Policy Learning @@ -22,7 +27,7 @@ python rl_scripts/run_XXX.py --task YYY --seed Z python rl_scripts/run_rombrl.py --task betan --seed 0 ``` - - XXX can be one of [rombrl, cql, edac, combo, mobile, bambrl, rambo], corresponding to our algorithm and the 6 baselines used in the paper. + - XXX can be one of [rombrl, cql, edac, combo, mobile, bamcts, rambo], corresponding to our algorithm and the 6 baselines used in the paper. - YYY can be one of [betan, dens_component1, rotation_component1], corresponding to the tracking tasks for betan, density, and rotation, respectively. - The random seed Z can be one of [0, 1, 2]. diff --git a/Fusion/offlinerlkit/policy/__init__.py b/Fusion/offlinerlkit/policy/__init__.py index e3c2829..dd60111 100644 --- a/Fusion/offlinerlkit/policy/__init__.py +++ b/Fusion/offlinerlkit/policy/__init__.py @@ -15,7 +15,7 @@ from offlinerlkit.policy.model_based.mobile import MOBILEPolicy from offlinerlkit.policy.model_based.rambo import RAMBOPolicy from offlinerlkit.policy.model_based.combo import COMBOPolicy -from offlinerlkit.policy.model_based.bambrl import BAMBRLPolicy +from offlinerlkit.policy.model_based.bamcts import BAMCTSPolicy from offlinerlkit.policy.model_based.rombrl import ROMBRLPolicy from offlinerlkit.policy.model_based.rombrl2 import ROMBRL2Policy @@ -33,7 +33,7 @@ "MOBILEPolicy", "RAMBOPolicy", "COMBOPolicy", - "BAMBRLPolicy", + "BAMCTSPolicy", "ROMBRLPolicy", "ROMBRL2Policy" ] \ No newline at end of file diff --git a/Fusion/offlinerlkit/policy/model_based/bambrl.py b/Fusion/offlinerlkit/policy/model_based/bamcts.py similarity index 99% rename from Fusion/offlinerlkit/policy/model_based/bambrl.py rename to Fusion/offlinerlkit/policy/model_based/bamcts.py index d050480..5f39cfb 100644 --- a/Fusion/offlinerlkit/policy/model_based/bambrl.py +++ b/Fusion/offlinerlkit/policy/model_based/bamcts.py @@ -10,7 +10,7 @@ from offlinerlkit.buffer import SLReplayBuffer, SL_Transition from torch.distributions import Normal, Independent -class BAMBRLPolicy(MOBILEPolicy): +class BAMCTSPolicy(MOBILEPolicy): def __init__( self, @@ -187,7 +187,7 @@ def rollout( step_actions = self.sa_processor.get_step_action(actions) full_actions[:, self.action_idxs] = step_actions.copy() - # new for mobile/bambrl + # new for mobile/bamcts rollout_transitions["hidden_states"].append(self.dynamics.get_memory()[idx_list]) next_observations, rewards, terminals, info = self.dynamics.step(priors, full_observations, pre_actions, full_actions, time_steps, time_terminals, self.state_idxs, init_samples['batch_idx_list'][t]) @@ -206,7 +206,7 @@ def rollout( rollout_transitions["rewards"].append(rewards[idx_list]) rollout_transitions["terminals"].append(terminals[idx_list]) - # new for mobile/bambrl + # new for mobile/bamcts rollout_transitions["full_obss"].append(full_observations[idx_list]) rollout_transitions["full_actions"].append(full_actions[idx_list]) rollout_transitions["pre_actions"].append(pre_actions[idx_list]) diff --git a/Fusion/requirements.txt b/Fusion/requirements.txt new file mode 100644 index 0000000..d475574 --- /dev/null +++ b/Fusion/requirements.txt @@ -0,0 +1,30 @@ +# Dependencies for the Tokamak Control experiments (Table 3). +# Versions are inferred from the code's imports and known-compatible combinations; +# if you hit a version conflict in your own environment, `pip freeze` after a +# successful install and pin exact versions here. +# +# This does NOT include `dynamics-toolbox`, which is a separate GitHub repo that +# must be cloned and installed with `pip install -e .` — see Fusion/README.md. +# DIII-D operational data required to run these experiments is proprietary and +# is not distributed with this repository (see Fusion/README.md). + +torch>=1.13,<2.1 +pytorch-lightning>=1.9,<2.0 +numpy>=1.21,<1.24 +scipy>=1.7 +pandas>=1.3 +matplotlib>=3.5 +scikit-learn>=1.0 +h5py>=3.6 +easydict +tqdm +Cython>=0.29,<3.0 +PyYAML>=6.0 + +# Config / experiment management +hydra-core>=1.2 +omegaconf>=2.2 +ray>=2.0 + +# RL environment +gym==0.23.1 diff --git a/Fusion/rl_scripts/run_bambrl.py b/Fusion/rl_scripts/run_bamcts.py similarity index 98% rename from Fusion/rl_scripts/run_bambrl.py rename to Fusion/rl_scripts/run_bamcts.py index 035ad54..226ef82 100644 --- a/Fusion/rl_scripts/run_bambrl.py +++ b/Fusion/rl_scripts/run_bamcts.py @@ -14,14 +14,14 @@ from offlinerlkit.buffer import BayesReplayBuffer, SLReplayBuffer from offlinerlkit.utils.logger import Logger, make_log_dirs from offlinerlkit.policy_trainer import MBPolicyTrainer -from offlinerlkit.policy import BAMBRLPolicy +from offlinerlkit.policy import BAMCTSPolicy from offlinerlkit.utils.searcher import Searcher from rl_preparation.get_rl_data_envs import get_rl_data_envs def get_args(): parser = argparse.ArgumentParser() - parser.add_argument("--algo-name", type=str, default="bambrl") + parser.add_argument("--algo-name", type=str, default="bamcts") parser.add_argument("--actor-lr", type=float, default=1e-4) parser.add_argument("--critic-lr", type=float, default=3e-4) parser.add_argument("--hidden-dims", type=int, nargs='*', default=[256, 256]) @@ -216,7 +216,7 @@ def train(args=get_args()): searcher = None # create policy - policy = BAMBRLPolicy( + policy = BAMCTSPolicy( args.use_ba, args.use_search, args.search_ratio, diff --git a/README.md b/README.md index dfe6062..f71b1f8 100644 --- a/README.md +++ b/README.md @@ -1,3 +1,83 @@ -# Policy-Driven World Model Adaptation for Robust Offline Model-based Reinforcement Learning +# ROMBRL: Policy-Driven World Model Adaptation for Robust Offline Model-based RL -Please refer to 'D4RL' and 'Fusion' for the results in Tables 1 and 3, respectively. The D4RL folder also includes instructions for the RWRL-style robustness evaluation in Table 4. Detailed README files are provided in each folder. +[![arXiv](https://img.shields.io/badge/arXiv-2505.13709-b31b1b.svg)](https://arxiv.org/abs/2505.13709) +[![OpenReview](https://img.shields.io/badge/OpenReview-forum-8c1b13.svg)](https://openreview.net/forum?id=eB5i7caric) +[![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](LICENSE) +[![ICML 2026](https://img.shields.io/badge/ICML-2026-blue.svg)](https://arxiv.org/abs/2505.13709) + +Official implementation for **"Policy-Driven World Model Adaptation for Robust Offline Model-based Reinforcement Learning"** (ICML 2026), by Jiayu Chen, Le Xu, Aravind Venugopal, and Jeff Schneider. + +- 📄 Paper: [arXiv:2505.13709](https://arxiv.org/abs/2505.13709) · [OpenReview](https://openreview.net/forum?id=eB5i7caric) +- 🌐 Project page: [agentic-intelligence-lab.github.io/ROMBRL](https://agentic-intelligence-lab.github.io/ROMBRL/) +- 📝 Blog post: [agentic-intelligence-lab.org/blog](https://agentic-intelligence-lab.org/blog/) +- 🖼️ Poster: [docs/assets/poster.pdf](docs/assets/poster.pdf) + +## TL;DR + +Offline model-based RL (MBRL) usually learns a world model and a policy in two separate stages — fit the model to maximize data likelihood, then optimize the policy against the fixed model. This objective mismatch leaves policies brittle to deployment-time noise. **ROMBRL** jointly adapts the world model *with* the policy under a single constrained maximin objective, solved via **Stackelberg learning dynamics** (policy as leader, world model as adversarial follower), with a formal suboptimality bound. It achieves state-of-the-art robustness on D4RL MuJoCo and stochastic Tokamak Control benchmarks, at almost no cost to clean-environment performance. + +## Repository Structure + +``` +ROMBRL/ +├── D4RL/ # D4RL MuJoCo experiments — Tables 1, 2, and 4 +└── Fusion/ # Tokamak Control experiments — Table 3 +``` + +Each folder is a self-contained fork of [OfflineRL-Kit](https://github.com/yihaosun1124/OfflineRL-Kit) extended with our method (`rombrl2`/`rombrl` policies) and baselines. See `D4RL/README.md` and `Fusion/README.md` for folder-specific setup and usage details. + +## Data Availability + +- **D4RL (Tables 1, 2, 4):** uses the public [D4RL](https://github.com/Farama-Foundation/d4rl) MuJoCo datasets — fully reproducible. +- **Tokamak Control (Table 3):** uses operational data from the DIII-D tokamak, which is **proprietary and not redistributed in this repository**. We are unable to release it until we obtain the necessary approvals; see `Fusion/README.md`. The `Fusion/` code (dynamics model, environment, RL pipeline) is provided for reference and can be run once you have access to equivalent data. + +## Installation + +Each experiment folder has its own dependencies (see `D4RL/requirements.txt` / `Fusion/requirements.txt`) since the D4RL and Tokamak Control setups use different environment and simulator stacks. We recommend separate virtual environments for each: + +```bash +# D4RL MuJoCo experiments +cd D4RL && pip install -r requirements.txt + +# Tokamak Control experiments +cd Fusion && pip install -r requirements.txt +``` + +Follow the full setup steps (MuJoCo binary, `ctree` extension build, `dynamics-toolbox`) in each folder's README before running experiments. + +**Docker (experimental):** `D4RL/Dockerfile` and `Fusion/Dockerfile` package each environment. They have not been build-tested against a live Docker daemon yet — please build and smoke-test locally before relying on them: + +```bash +cd D4RL && docker build -t rombrl-d4rl . && docker run --rm -it rombrl-d4rl +cd Fusion && docker build -t rombrl-fusion . && docker run --rm -it -v /path/to/data:/workspace/Fusion/data rombrl-fusion +``` + +## Reproducing the Paper's Results + +| Paper Table | Experiment | Folder | Example command | +|---|---|---|---| +| Table 1 | D4RL MuJoCo under 5% measurement noise | `D4RL/` | `python run_rombrl2.py` (edit `load_path_id` for task/seed) | +| Table 2 | Standard vs. noisy performance drop | `D4RL/` | same runs as Table 1, evaluated with/without `--noise_scale 0.05` | +| Table 3 | Tokamak Control tracking tasks | `Fusion/` | `python rl_scripts/run_rombrl.py --task betan --seed 0` | +| Table 4 | RWRL-style deployment perturbations | `D4RL/` | `python run_rombrl2.py --task halfcheetah-medium-replay-v2 --challenge sensor_dropped --noise_scale 0.0` | + +## Citation + +The official PMLR proceedings for ICML 2026 (volume 306) have not been posted yet; please cite the arXiv version for now: + +```bibtex +@article{chen2025rombrl, + title = {Policy-Driven World Model Adaptation for Robust Offline Model-based Reinforcement Learning}, + author = {Chen, Jiayu and Xu, Le and Venugopal, Aravind and Schneider, Jeff}, + journal = {arXiv preprint arXiv:2505.13709}, + year = {2025} +} +``` + +## Acknowledgements + +This work was funded in part by the Department of Energy Fusion Energy Sciences under grant DE-SC0024544. The codebase builds on [OfflineRL-Kit](https://github.com/yihaosun1124/OfflineRL-Kit). + +## License + +Released under the [MIT License](LICENSE). diff --git a/docs/assets/cmu_logo.png b/docs/assets/cmu_logo.png new file mode 100644 index 0000000..e154891 Binary files /dev/null and b/docs/assets/cmu_logo.png differ diff --git a/docs/assets/fig1_motivation.png b/docs/assets/fig1_motivation.png new file mode 100644 index 0000000..1756d44 Binary files /dev/null and b/docs/assets/fig1_motivation.png differ diff --git a/docs/assets/fig3_ablation.png b/docs/assets/fig3_ablation.png new file mode 100644 index 0000000..7af3ba3 Binary files /dev/null and b/docs/assets/fig3_ablation.png differ diff --git a/docs/assets/hku_logo.png b/docs/assets/hku_logo.png new file mode 100644 index 0000000..b62db8c Binary files /dev/null and b/docs/assets/hku_logo.png differ diff --git a/docs/assets/infiforce_logo.png b/docs/assets/infiforce_logo.png new file mode 100644 index 0000000..9c6c5eb Binary files /dev/null and b/docs/assets/infiforce_logo.png differ diff --git a/docs/assets/poster.pdf b/docs/assets/poster.pdf new file mode 100644 index 0000000..f0ba547 Binary files /dev/null and b/docs/assets/poster.pdf differ diff --git a/docs/assets/poster.png b/docs/assets/poster.png new file mode 100644 index 0000000..bf116f7 Binary files /dev/null and b/docs/assets/poster.png differ diff --git a/docs/assets/thu_logo.png b/docs/assets/thu_logo.png new file mode 100644 index 0000000..7f182b5 Binary files /dev/null and b/docs/assets/thu_logo.png differ diff --git a/docs/assets/tokamak_illustration.png b/docs/assets/tokamak_illustration.png new file mode 100644 index 0000000..b641fd1 Binary files /dev/null and b/docs/assets/tokamak_illustration.png differ diff --git a/docs/index.html b/docs/index.html new file mode 100644 index 0000000..b80f06a --- /dev/null +++ b/docs/index.html @@ -0,0 +1,277 @@ + + + + + +ROMBRL: Policy-Driven World Model Adaptation for Robust Offline Model-based RL + + + + + + + +
+ +
+ ICML 2026 · Seoul, South Korea +

Policy-Driven World Model Adaptation for Robust Offline Model-based Reinforcement Learning

+
+ Jiayu Chen1,2,*  ·  + Le Xu3,*  ·  + Aravind Venugopal4  ·  + Jeff Schneider4 +
+
+ 1The University of Hong Kong   2INFIFORCE Intelligent Technology Co., Ltd.   + 3Tsinghua University   4Carnegie Mellon University   *Equal contribution +
+ +
+ The University of Hong Kong + Tsinghua University + Carnegie Mellon University + Infiforce Intelligent Technology +
+ + +
+ +

Abstract

+

+ Offline reinforcement learning (RL) offers a powerful paradigm for data-driven control. Compared to model-free approaches, offline model-based RL (MBRL) explicitly learns a world model from a static dataset and uses it as a surrogate simulator, improving data efficiency and enabling potential generalization beyond the dataset support. However, most existing offline MBRL methods follow a two-stage training procedure: first learning a world model by maximizing the likelihood of the observed transitions, then optimizing a policy to maximize its expected return under the learned model. This objective mismatch results in a world model that is not necessarily optimized for effective policy learning. Moreover, we observe that policies learned via offline MBRL often lack robustness during deployment, and small adversarial noise in the environment can lead to significant performance degradation. To address these, we propose a framework that dynamically adapts the world model alongside the policy under a unified learning objective aimed at improving robustness. At the core of our method is a maximin optimization problem, which we solve by innovatively utilizing Stackelberg learning dynamics. We provide theoretical analysis to support our design and introduce computationally efficient implementations. We benchmark our algorithm on twelve noisy D4RL MuJoCo tasks and three stochastic Tokamak Control tasks, demonstrating its state-of-the-art performance. +

+ +

Motivation

+

+ Existing offline MBRL methods learn a world model to maximize data likelihood, then freeze it (or lightly adapt it) while optimizing the policy. This causes an objective mismatch — the model is trained to explain the offline data, not to support good policy learning. In practice, this leaves policies brittle: even a modest amount of measurement noise at deployment time can sharply degrade the performance of strong offline RL baselines. +

+
+ Average performance of offline RL algorithms before and after deployment noise +
Average scores of offline RL algorithms on nine D4RL MuJoCo tasks, before and after injecting 5% measurement noise into state transitions at deployment. Most methods — including the model-based robust baseline RAMBO — lose a substantial fraction of their performance.
+
+ +

Method

+

+ We formulate offline MBRL as a constrained maximin problem: the policy maximizes its return under the worst-case world model within an uncertainty set consistent with the offline data, while the world model is adversarially updated to minimize that return within a trust region. +

+
\[ \max_{\theta} J(\theta, \phi') \quad \text{s.t.} \quad \phi' \in \operatorname*{arg\,min}_{\phi \in \Phi} J(\theta, \phi) \]
+
\[ \Phi = \left\{\, \phi \in \mathcal{M} \;:\; \mathbb{E}_{(s,a)\sim\mathcal{D}}\Big[\mathrm{KL}\big(P_{\hat\phi}(\cdot\mid s,a) \,\|\, P_\phi(\cdot\mid s,a)\big)\Big] \le \epsilon \,\right\} \]
+

+ We solve this as a Stackelberg game — the policy is the leader, the world model is the follower that best-responds to the policy. This is the opposite direction from online MBRL, where the model is typically adapted to serve the leader's interest rather than oppose it. Using implicit differentiation through the follower's best response, we derive primal-dual Stackelberg update rules for \((\theta, \phi, \lambda)\), made practical at scale by: +

+ +

The full algorithm, ROMBRL, is given in Appendix J of the paper.

+ +

Theory

+

+ Theorem 3.1 bounds the performance gap between the policy learned by ROMBRL and the best possible comparator policy, assuming the true environment lies in the uncertainty set \(\Phi\) with probability at least \(1 - \delta/2\): +

+
\[ J(\theta^*, \phi^*) - J(\hat\theta, \phi^*) \;\le\; \frac{\sqrt{C}}{(1-\gamma)^2} \sqrt{\, 4\epsilon + c\left(\sqrt{\frac{\log(2|\Phi|/\delta)}{N}} + \frac{\log(2|\Phi|/\delta)}{N}\right) } \]
+

+ where \(N\) is the offline dataset size, \(|\Phi|\) the covering number of the uncertainty set, \(C\) a concentrability coefficient, and \(\epsilon\) the uncertainty-set radius — so the suboptimality gap shrinks as the dataset grows, giving a formal robustness guarantee rather than a purely heuristic one. Theorems 3.2 and 3.3 instantiate this bound for tabular MDPs and for continuous MDPs with Gaussian world models respectively, characterizing how \(\epsilon\) should scale with dataset size and dimensionality. +

+ +

Results

+

D4RL MuJoCo, under 5% measurement noise at deployment. ROMBRL ranks first on 7/12 tasks and second on 4, with the best average score by a wide margin (Cohen's d ≥ 2 over every baseline).

+ + + + +
MethodROMBRL (ours)CQLEDACCOMBORAMBOMOBILERORLTRACERRFQI
Average Score77.7 (0.5)60.7 (1.2)53.7 (4.6)55.5 (3.6)55.8 (1.3)70.7 (2.4)62.3 (0.2)44.1 (4.8)30.9 (2.1)
Cohen's d vs. ROMBRL–18.97.48.522.94.140.49.830.7
+ +

Robustness vs. clean-environment trade-off (OfflineRL-Kit protocol). ROMBRL matches the best clean-environment performance while losing almost nothing to noise.

+ + + + + +
MetricROMBRL (ours)CQLEDACCOMBORAMBOMOBILERORL
Standard env.92.880.493.089.382.795.989.5
Noisy env.93.477.568.372.266.085.378.7
Performance drop ↓-0.6%3.6%26.6%19.1%20.2%11.1%12.1%
+ +

Tokamak Control (negative tracking error; higher is better). ROMBRL ranks first on 2/3 targets and second on the third, with the lowest variance across seeds.

+
+ ROMBRL applied to Tokamak plasma control +
ROMBRL applied to Tokamak Control: an RL controller trained on a surrogate dynamics model to drive plasma profiles toward a target via actuators such as power, torque, and ECH.
+
+ + + + + + +
Tracking TargetROMBRL (ours)CQLEDACCOMBORAMBOMOBILEBAMCTS
βN-70.9* (0.9)-78.4 (3.1)-63.4 (1.7)-84.3 (7.6)-121.1 (19.9)-133.9 (10.1)-111.3 (24.3)
Density-60.0 (1.9)-87.3 (12.5)-112.5 (11.1)-67.0* (3.1)-81.3 (15.7)-75.3 (4.3)-79.6 (13.8)
Rotation-10.6 (3.7)-39.2* (10.1)-95.4 (64.3)-69.6 (25.9)-300.3 (260.5)-257.6 (153.7)-305.6 (242.6)
Average Return-47.1 (1.2)-68.3* (6.8)-90.4 (11.5)-73.6 (5.8)-167.6 (91.6)-155.5 (47.7)-165.5 (84.5)
+

* marks the second-best result for each row.

+ +

RWRL-style deployment perturbations (dropped/stuck observations, body mass, friction, joint damping). ROMBRL achieves the best score on every perturbation type and the smallest overall performance drop.

+ + + + +
MetricROMBRL (ours)CQLMOBILERAMBORORLTRACERRFQI
Average score under perturbation51.9 (1.0)31.2 (1.0)45.7 (1.9)46.2 (0.2)39.3 (1.8)23.4 (1.6)16.3 (1.6)
Performance drop ↓25.4%33.2%31.6%30.0%37.6%36.2%52.1%
+ +
+ Ablation of Stackelberg gradient update mechanisms +
Ablation on walker2d-medium: the full constrained Stackelberg update (Ours) substantially outperforms a naive alternating update and an unconstrained Stackelberg update — robustness comes from anticipating how the uncertainty-set boundary shifts with the policy, not just from anticipating the model.
+
+ +

Poster

+ +
+ ROMBRL ICML 2026 poster +
Click for the full-resolution poster PDF.
+
+
+ +

Citation

+
The official PMLR proceedings for ICML 2026 (volume 306) have not been posted yet — please cite the arXiv version for now.
+
@article{chen2025rombrl,
+  title   = {Policy-Driven World Model Adaptation for Robust Offline Model-based Reinforcement Learning},
+  author  = {Chen, Jiayu and Xu, Le and Venugopal, Aravind and Schneider, Jeff},
+  journal = {arXiv preprint arXiv:2505.13709},
+  year    = {2025}
+}
+ + + +
+ + +