Fixes build errors due to name conflicts

This commit is contained in:
cvoelcker
2025-07-21 18:31:20 -04:00
parent 094ee0c5ba
commit e2f99648ae
26 changed files with 52 additions and 79 deletions
View File
+95
View File
@@ -0,0 +1,95 @@
import torch
from omegaconf import DictConfig
def make_envs(cfg: DictConfig, device: torch.device, seed: int = None) -> tuple:
if cfg.env.type == "humanoid_bench":
from reppo_alg.env_utils.torch_wrappers.humanoid_bench_env import (
HumanoidBenchEnv,
)
envs = HumanoidBenchEnv(
cfg.env.name, cfg.hyperparameters.num_envs, device=device
)
return envs, envs
elif cfg.env.type == "isaaclab":
from reppo_alg.env_utils.torch_wrappers.isaaclab_env import IsaacLabEnv
envs = IsaacLabEnv(
cfg.env.name,
device.type,
cfg.hyperparameters.num_envs,
cfg=seed,
action_bounds=cfg.env.action_bounds,
)
return envs, envs
elif cfg.env.type == "mjx":
from reppo_alg.env_utils.torch_wrappers.mujoco_playground_env import make_env
# TODO: Check if re-using same envs for eval could reduce memory usage
envs, eval_envs = make_env(
env_name=cfg.env.name,
seed=seed,
num_envs=cfg.hyperparameters.num_envs,
num_eval_envs=cfg.hyperparameters.num_envs,
device_rank=cfg.platform.device_rank,
use_domain_randomization=False,
use_push_randomization=True,
)
return envs, eval_envs
elif cfg.env.type == "maniskill":
import gymnasium as gym
import mani_skill.envs # noqa: F401
from mani_skill.utils import gym_utils
from mani_skill.utils.wrappers.flatten import FlattenActionSpaceWrapper
from mani_skill.vector.wrappers.gymnasium import ManiSkillVectorEnv
from reppo_alg.env_utils.torch_wrappers.maniskill_wrapper import (
ManiSkillWrapper,
)
envs = gym.make(
cfg.env.name,
num_envs=cfg.hyperparameters.num_envs,
reconfiguration_freq=None,
**cfg.env.env_kwargs,
)
eval_envs = gym.make(
cfg.env.name,
num_envs=cfg.hyperparameters.num_envs,
reconfiguration_freq=1,
**cfg.env.env_kwargs,
)
cfg.env.max_episode_steps = gym_utils.find_max_episode_steps_value(envs)
# heuristic for setting gamma
cfg.hyperparameters.gamma = 1.0 - 10.0 / cfg.env.max_episode_steps
if isinstance(envs.action_space, gym.spaces.Dict):
envs = FlattenActionSpaceWrapper(envs)
eval_envs = FlattenActionSpaceWrapper(eval_envs)
envs = ManiSkillVectorEnv(
envs,
cfg.hyperparameters.num_envs,
ignore_terminations=not cfg.env.partial_reset,
record_metrics=True,
)
eval_envs = ManiSkillVectorEnv(
eval_envs,
cfg.hyperparameters.num_envs,
ignore_terminations=True,
record_metrics=True,
)
return ManiSkillWrapper(
envs,
max_episode_steps=cfg.env.max_episode_steps,
partial_reset=cfg.env.partial_reset,
device=device.type,
), ManiSkillWrapper(
eval_envs,
max_episode_steps=cfg.env.max_episode_steps,
partial_reset=cfg.env.partial_reset,
device=device.type,
)
else:
raise ValueError(
f"Unknown environment type: {cfg.env.type}. Supported types are 'humanoid_bench', 'isaaclab', 'maniskill', and 'mjx'."
)
+691
View File
@@ -0,0 +1,691 @@
import os
import sys
os.environ["TORCHDYNAMO_INLINE_INBUILT_NN_MODULES"] = "1"
os.environ["OMP_NUM_THREADS"] = "1"
if sys.platform != "darwin":
os.environ["MUJOCO_GL"] = "egl"
else:
os.environ["MUJOCO_GL"] = "glfw"
os.environ["XLA_PYTHON_CLIENT_PREALLOCATE"] = "false"
os.environ["JAX_DEFAULT_MATMUL_PRECISION"] = "highest"
import math
import random
import time
import numpy as np
import tqdm
import wandb
try:
# Required for avoiding IsaacGym import error
import isaacgym
except ImportError:
pass
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from reppo_alg.torchrl.reppo import (
EmpiricalNormalization,
PerTaskRewardNormalizer,
RewardNormalizer,
SimpleReplayBuffer,
save_params,
)
from hyperparams import get_args
from tensordict import TensorDict
from torch.amp import GradScaler, autocast
torch.set_float32_matmul_precision("high")
def main():
args = get_args()
print(args)
run_name = f"{args.env_name}__{args.exp_name}__{args.seed}"
amp_enabled = args.amp and args.cuda and torch.cuda.is_available()
amp_device_type = (
"cuda"
if args.cuda and torch.cuda.is_available()
else "mps"
if args.cuda and torch.backends.mps.is_available()
else "cpu"
)
amp_dtype = torch.bfloat16 if args.amp_dtype == "bf16" else torch.float16
scaler = GradScaler(enabled=amp_enabled and amp_dtype == torch.float16)
if args.use_wandb:
wandb.init(
project=args.project,
name=run_name,
config=vars(args),
save_code=True,
)
random.seed(args.seed)
np.random.seed(args.seed)
torch.manual_seed(args.seed)
torch.backends.cudnn.deterministic = args.torch_deterministic
if not args.cuda:
device = torch.device("cpu")
else:
if torch.cuda.is_available():
device = torch.device(f"cuda:{args.device_rank}")
elif torch.backends.mps.is_available():
device = torch.device(f"mps:{args.device_rank}")
else:
raise ValueError("No GPU available")
print(f"Using device: {device}")
if args.env_name.startswith("h1hand-") or args.env_name.startswith("h1-"):
from reppo_alg.env_utils.torch_wrappers.humanoid_bench_env import (
HumanoidBenchEnv,
)
env_type = "humanoid_bench"
envs = HumanoidBenchEnv(args.env_name, args.num_envs, device=device)
eval_envs = envs
elif args.env_name.startswith("Isaac-"):
from reppo_alg.env_utils.torch_wrappers.isaaclab_env import IsaacLabEnv
env_type = "isaaclab"
envs = IsaacLabEnv(
args.env_name,
device.type,
args.num_envs,
args.seed,
action_bounds=args.action_bounds,
)
eval_envs = envs
elif args.env_name.startswith("MTBench-"):
from reppo_alg.env_utils.torch_wrappers.mtbench_env import MTBenchEnv
env_name = "-".join(args.env_name.split("-")[1:])
env_type = "mtbench"
envs = MTBenchEnv(env_name, args.device_rank, args.num_envs, args.seed)
eval_envs = envs
else:
from reppo_alg.env_utils.torch_wrappers.mujoco_playground_env import make_env
# TODO: Check if re-using same envs for eval could reduce memory usage
env_type = "mujoco_playground"
envs, eval_envs = make_env(
args.env_name,
args.seed,
args.num_envs,
args.num_eval_envs,
args.device_rank,
use_tuned_reward=args.use_tuned_reward,
use_domain_randomization=args.use_domain_randomization,
use_push_randomization=args.use_push_randomization,
)
n_act = envs.num_actions
n_obs = envs.num_obs if isinstance(envs.num_obs, int) else envs.num_obs[0]
if envs.asymmetric_obs:
n_critic_obs = (
envs.num_privileged_obs
if isinstance(envs.num_privileged_obs, int)
else envs.num_privileged_obs[0]
)
else:
n_critic_obs = n_obs
action_low, action_high = -1.0, 1.0
if args.obs_normalization:
obs_normalizer = EmpiricalNormalization(shape=n_obs, device=device)
critic_obs_normalizer = EmpiricalNormalization(
shape=n_critic_obs, device=device
)
else:
obs_normalizer = nn.Identity()
critic_obs_normalizer = nn.Identity()
if args.reward_normalization:
if env_type in ["mtbench"]:
reward_normalizer = PerTaskRewardNormalizer(
num_tasks=envs.num_tasks,
gamma=args.gamma,
device=device,
g_max=min(abs(args.v_min), abs(args.v_max)),
)
else:
reward_normalizer = RewardNormalizer(
gamma=args.gamma,
device=device,
g_max=min(abs(args.v_min), abs(args.v_max)),
)
else:
reward_normalizer = nn.Identity()
actor_kwargs = {
"n_obs": n_obs,
"n_act": n_act,
"num_envs": args.num_envs,
"device": device,
"init_scale": args.init_scale,
"hidden_dim": args.actor_hidden_dim,
}
critic_kwargs = {
"n_obs": n_critic_obs,
"n_act": n_act,
"num_atoms": args.num_atoms,
"v_min": args.v_min,
"v_max": args.v_max,
"hidden_dim": args.critic_hidden_dim,
"device": device,
}
if env_type == "mtbench":
actor_kwargs["n_obs"] = n_obs - envs.num_tasks + args.task_embedding_dim
critic_kwargs["n_obs"] = n_critic_obs - envs.num_tasks + args.task_embedding_dim
actor_kwargs["num_tasks"] = envs.num_tasks
actor_kwargs["task_embedding_dim"] = args.task_embedding_dim
critic_kwargs["num_tasks"] = envs.num_tasks
critic_kwargs["task_embedding_dim"] = args.task_embedding_dim
if args.agent == "fasttd3":
if env_type in ["mtbench"]:
from reppo_alg.network_utils.fast_td3_nets import (
MultiTaskActor,
MultiTaskCritic,
)
actor_cls = MultiTaskActor
critic_cls = MultiTaskCritic
else:
from reppo_alg.network_utils.fast_td3_nets import Actor, Critic
actor_cls = Actor
critic_cls = Critic
print("Using FastTD3")
elif args.agent == "fasttd3_simbav2":
if env_type in ["mtbench"]:
from reppo_alg.network_utils.fast_td3_nets_simbav2 import (
MultiTaskActor,
MultiTaskCritic,
)
actor_cls = MultiTaskActor
critic_cls = MultiTaskCritic
else:
from reppo_alg.network_utils.fast_td3_nets_simbav2 import Actor, Critic
actor_cls = Actor
critic_cls = Critic
print("Using FastTD3 + SimbaV2")
actor_kwargs.pop("init_scale")
actor_kwargs.update(
{
"scaler_init": math.sqrt(2.0 / args.actor_hidden_dim),
"scaler_scale": math.sqrt(2.0 / args.actor_hidden_dim),
"alpha_init": 1.0 / (args.actor_num_blocks + 1),
"alpha_scale": 1.0 / math.sqrt(args.actor_hidden_dim),
"expansion": 4,
"c_shift": 3.0,
"num_blocks": args.actor_num_blocks,
}
)
critic_kwargs.update(
{
"scaler_init": math.sqrt(2.0 / args.critic_hidden_dim),
"scaler_scale": math.sqrt(2.0 / args.critic_hidden_dim),
"alpha_init": 1.0 / (args.critic_num_blocks + 1),
"alpha_scale": 1.0 / math.sqrt(args.critic_hidden_dim),
"num_blocks": args.critic_num_blocks,
"expansion": 4,
"c_shift": 3.0,
}
)
else:
raise ValueError(f"Agent {args.agent} not supported")
actor = actor_cls(**actor_kwargs)
if env_type in ["mtbench"]:
# Python 3.8 doesn't support 'from_module' in tensordict
policy = actor.explore
else:
from tensordict import from_module
actor_detach = actor_cls(**actor_kwargs)
# Copy params to actor_detach without grad
from_module(actor).data.to_module(actor_detach)
policy = actor_detach.explore
qnet = critic_cls(**critic_kwargs)
qnet_target = critic_cls(**critic_kwargs)
qnet_target.load_state_dict(qnet.state_dict())
q_optimizer = optim.AdamW(
list(qnet.parameters()),
lr=torch.tensor(args.critic_learning_rate, device=device),
weight_decay=args.weight_decay,
)
actor_optimizer = optim.AdamW(
list(actor.parameters()),
lr=torch.tensor(args.actor_learning_rate, device=device),
weight_decay=args.weight_decay,
)
# Add learning rate schedulers
q_scheduler = optim.lr_scheduler.CosineAnnealingLR(
q_optimizer,
T_max=args.total_timesteps,
eta_min=torch.tensor(args.critic_learning_rate_end, device=device),
)
actor_scheduler = optim.lr_scheduler.CosineAnnealingLR(
actor_optimizer,
T_max=args.total_timesteps,
eta_min=torch.tensor(args.actor_learning_rate_end, device=device),
)
rb = SimpleReplayBuffer(
n_env=args.num_envs,
buffer_size=args.buffer_size,
n_obs=n_obs,
n_act=n_act,
n_critic_obs=n_critic_obs,
asymmetric_obs=envs.asymmetric_obs,
playground_mode=env_type == "mujoco_playground",
n_steps=args.num_steps,
gamma=args.gamma,
device=device,
)
policy_noise = args.policy_noise
noise_clip = args.noise_clip
def evaluate():
obs_normalizer.eval()
num_eval_envs = eval_envs.num_envs
episode_returns = torch.zeros(num_eval_envs, device=device)
episode_lengths = torch.zeros(num_eval_envs, device=device)
done_masks = torch.zeros(num_eval_envs, dtype=torch.bool, device=device)
if env_type == "isaaclab":
obs = eval_envs.reset(random_start_init=False)
else:
obs = eval_envs.reset()
# Run for a fixed number of steps
for i in range(eval_envs.max_episode_steps):
with (
torch.no_grad(),
autocast(
device_type=amp_device_type, dtype=amp_dtype, enabled=amp_enabled
),
):
obs = normalize_obs(obs)
actions = actor(obs)
next_obs, rewards, dones, _, infos = eval_envs.step(actions.float())
if env_type == "mtbench":
# We only report success rate in MTBench evaluation
rewards = (
infos["episode"]["success"].float() if "episode" in infos else 0.0
)
episode_returns = torch.where(
~done_masks, episode_returns + rewards, episode_returns
)
episode_lengths = torch.where(
~done_masks, episode_lengths + 1, episode_lengths
)
if env_type == "mtbench" and "episode" in infos:
dones = dones | infos["episode"]["success"]
done_masks = torch.logical_or(done_masks, dones)
if done_masks.all():
break
obs = next_obs
obs_normalizer.train()
return episode_returns.mean().item(), episode_lengths.mean().item()
def update_main(data, logs_dict):
with autocast(
device_type=amp_device_type, dtype=amp_dtype, enabled=amp_enabled
):
observations = data["observations"]
next_observations = data["next"]["observations"]
if envs.asymmetric_obs:
critic_observations = data["critic_observations"]
next_critic_observations = data["next"]["critic_observations"]
else:
critic_observations = observations
next_critic_observations = next_observations
actions = data["actions"]
rewards = data["next"]["rewards"]
dones = data["next"]["dones"].bool()
truncations = data["next"]["truncations"].bool()
if args.disable_bootstrap:
bootstrap = (~dones).float()
else:
bootstrap = (truncations | ~dones).float()
clipped_noise = torch.randn_like(actions)
clipped_noise = clipped_noise.mul(policy_noise).clamp(
-noise_clip, noise_clip
)
next_state_actions = (actor(next_observations) + clipped_noise).clamp(
action_low, action_high
)
discount = args.gamma ** data["next"]["effective_n_steps"]
with torch.no_grad():
qf1_next_target_projected, qf2_next_target_projected = (
qnet_target.projection(
next_critic_observations,
next_state_actions,
rewards,
bootstrap,
discount,
)
)
qf1_next_target_value = qnet_target.get_value(qf1_next_target_projected)
qf2_next_target_value = qnet_target.get_value(qf2_next_target_projected)
if args.use_cdq:
qf_next_target_dist = torch.where(
qf1_next_target_value.unsqueeze(1)
< qf2_next_target_value.unsqueeze(1),
qf1_next_target_projected,
qf2_next_target_projected,
)
qf1_next_target_dist = qf2_next_target_dist = qf_next_target_dist
else:
qf1_next_target_dist, qf2_next_target_dist = (
qf1_next_target_projected,
qf2_next_target_projected,
)
qf1, qf2 = qnet(critic_observations, actions)
qf1_loss = -torch.sum(
qf1_next_target_dist * F.log_softmax(qf1, dim=1), dim=1
).mean()
qf2_loss = -torch.sum(
qf2_next_target_dist * F.log_softmax(qf2, dim=1), dim=1
).mean()
qf_loss = qf1_loss + qf2_loss
q_optimizer.zero_grad(set_to_none=True)
scaler.scale(qf_loss).backward()
scaler.unscale_(q_optimizer)
critic_grad_norm = torch.nn.utils.clip_grad_norm_(
qnet.parameters(),
max_norm=args.max_grad_norm if args.max_grad_norm > 0 else float("inf"),
)
scaler.step(q_optimizer)
scaler.update()
q_scheduler.step()
logs_dict["critic_grad_norm"] = critic_grad_norm.detach()
logs_dict["qf_loss"] = qf_loss.detach()
logs_dict["qf_max"] = qf1_next_target_value.max().detach()
logs_dict["qf_min"] = qf1_next_target_value.min().detach()
return logs_dict
def update_pol(data, logs_dict):
with autocast(
device_type=amp_device_type, dtype=amp_dtype, enabled=amp_enabled
):
critic_observations = (
data["critic_observations"]
if envs.asymmetric_obs
else data["observations"]
)
qf1, qf2 = qnet(critic_observations, actor(data["observations"]))
qf1_value = qnet.get_value(F.softmax(qf1, dim=1))
qf2_value = qnet.get_value(F.softmax(qf2, dim=1))
if args.use_cdq:
qf_value = torch.minimum(qf1_value, qf2_value)
else:
qf_value = (qf1_value + qf2_value) / 2.0
actor_loss = -qf_value.mean()
actor_optimizer.zero_grad(set_to_none=True)
scaler.scale(actor_loss).backward()
scaler.unscale_(actor_optimizer)
actor_grad_norm = torch.nn.utils.clip_grad_norm_(
actor.parameters(),
max_norm=args.max_grad_norm if args.max_grad_norm > 0 else float("inf"),
)
scaler.step(actor_optimizer)
scaler.update()
actor_scheduler.step()
logs_dict["actor_grad_norm"] = actor_grad_norm.detach()
logs_dict["actor_loss"] = actor_loss.detach()
return logs_dict
if args.compile:
mode = None
update_main = torch.compile(update_main, mode=mode)
update_pol = torch.compile(update_pol, mode=mode)
policy = torch.compile(policy, mode=mode)
normalize_obs = torch.compile(obs_normalizer.forward, mode=mode)
normalize_critic_obs = torch.compile(critic_obs_normalizer.forward, mode=mode)
if args.reward_normalization:
update_stats = torch.compile(reward_normalizer.update_stats, mode=mode)
normalize_reward = torch.compile(reward_normalizer.forward, mode=mode)
else:
normalize_obs = obs_normalizer.forward
normalize_critic_obs = critic_obs_normalizer.forward
if args.reward_normalization:
update_stats = reward_normalizer.update_stats
normalize_reward = reward_normalizer.forward
if envs.asymmetric_obs:
obs, critic_obs = envs.reset_with_critic_obs()
critic_obs = torch.as_tensor(critic_obs, device=device, dtype=torch.float)
else:
obs = envs.reset()
if args.checkpoint_path:
# Load checkpoint if specified
torch_checkpoint = torch.load(
f"{args.checkpoint_path}", map_location=device, weights_only=False
)
actor.load_state_dict(torch_checkpoint["actor_state_dict"])
obs_normalizer.load_state_dict(torch_checkpoint["obs_normalizer_state"])
critic_obs_normalizer.load_state_dict(
torch_checkpoint["critic_obs_normalizer_state"]
)
qnet.load_state_dict(torch_checkpoint["qnet_state_dict"])
qnet_target.load_state_dict(torch_checkpoint["qnet_target_state_dict"])
global_step = torch_checkpoint["global_step"]
else:
global_step = 0
dones = None
pbar = tqdm.tqdm(total=args.total_timesteps, initial=global_step)
start_time = None
desc = ""
while global_step < args.total_timesteps:
logs_dict = TensorDict()
if (
start_time is None
and global_step >= args.measure_burnin + args.learning_starts
):
start_time = time.time()
measure_burnin = global_step
with (
torch.no_grad(),
autocast(device_type=amp_device_type, dtype=amp_dtype, enabled=amp_enabled),
):
norm_obs = normalize_obs(obs)
actions = policy(obs=norm_obs, dones=dones)
next_obs, rewards, dones, _, infos = envs.step(actions.float())
print(infos["time_outs"])
truncations = infos["time_outs"]
if args.reward_normalization:
if env_type == "mtbench":
task_ids_one_hot = obs[..., -envs.num_tasks :]
task_indices = torch.argmax(task_ids_one_hot, dim=1)
update_stats(rewards, dones.float(), task_ids=task_indices)
else:
update_stats(rewards, dones.float())
if envs.asymmetric_obs:
next_critic_obs = infos["observations"]["critic"]
# Compute 'true' next_obs and next_critic_obs for saving
true_next_obs = torch.where(
dones[:, None] > 0, infos["observations"]["raw"]["obs"], next_obs
)
if envs.asymmetric_obs:
true_next_critic_obs = torch.where(
dones[:, None] > 0,
infos["observations"]["raw"]["critic_obs"],
next_critic_obs,
)
transition = TensorDict(
{
"observations": obs,
"actions": torch.as_tensor(actions, device=device, dtype=torch.float),
"next": {
"observations": true_next_obs,
"rewards": torch.as_tensor(
rewards, device=device, dtype=torch.float
),
"truncations": truncations.long(),
"dones": dones.long(),
},
},
batch_size=(envs.num_envs,),
device=device,
)
if envs.asymmetric_obs:
transition["critic_observations"] = critic_obs
transition["next"]["critic_observations"] = true_next_critic_obs
obs = next_obs
if envs.asymmetric_obs:
critic_obs = next_critic_obs
rb.extend(transition)
batch_size = args.batch_size // args.num_envs
if global_step > args.learning_starts:
for i in range(args.num_updates):
data = rb.sample(batch_size)
data["observations"] = normalize_obs(data["observations"])
data["next"]["observations"] = normalize_obs(
data["next"]["observations"]
)
raw_rewards = data["next"]["rewards"]
if env_type in ["mtbench"] and args.reward_normalization:
# Multi-task reward normalization
task_ids_one_hot = data["observations"][..., -envs.num_tasks :]
task_indices = torch.argmax(task_ids_one_hot, dim=1)
data["next"]["rewards"] = normalize_reward(
raw_rewards, task_ids=task_indices
)
else:
data["next"]["rewards"] = normalize_reward(raw_rewards)
if envs.asymmetric_obs:
data["critic_observations"] = normalize_critic_obs(
data["critic_observations"]
)
data["next"]["critic_observations"] = normalize_critic_obs(
data["next"]["critic_observations"]
)
logs_dict = update_main(data, logs_dict)
if args.num_updates > 1:
if i % args.policy_frequency == 1:
logs_dict = update_pol(data, logs_dict)
else:
if global_step % args.policy_frequency == 0:
logs_dict = update_pol(data, logs_dict)
for param, target_param in zip(
qnet.parameters(), qnet_target.parameters()
):
target_param.data.copy_(
args.tau * param.data + (1 - args.tau) * target_param.data
)
if global_step % 100 == 0 and start_time is not None:
speed = (global_step - measure_burnin) / (time.time() - start_time)
pbar.set_description(f"{speed: 4.4f} sps, " + desc)
with torch.no_grad():
logs = {
"actor_loss": logs_dict["actor_loss"].mean(),
"qf_loss": logs_dict["qf_loss"].mean(),
"qf_max": logs_dict["qf_max"].mean(),
"qf_min": logs_dict["qf_min"].mean(),
"actor_grad_norm": logs_dict["actor_grad_norm"].mean(),
"critic_grad_norm": logs_dict["critic_grad_norm"].mean(),
"env_rewards": rewards.mean(),
"buffer_rewards": raw_rewards.mean(),
}
if args.eval_interval > 0 and global_step % args.eval_interval == 0:
print(f"Evaluating at global step {global_step}")
eval_avg_return, eval_avg_length = evaluate()
if env_type in ["humanoid_bench", "isaaclab", "mtbench"]:
# NOTE: Hacky way of evaluating performance, but just works
obs = envs.reset()
logs["eval_avg_return"] = eval_avg_return
logs["eval_avg_length"] = eval_avg_length
if args.use_wandb:
wandb.log(
{
"speed": speed,
"frame": global_step * args.num_envs,
"critic_lr": q_scheduler.get_last_lr()[0],
"actor_lr": actor_scheduler.get_last_lr()[0],
**logs,
},
step=global_step,
)
if (
args.save_interval > 0
and global_step > 0
and global_step % args.save_interval == 0
):
print(f"Saving model at global step {global_step}")
save_params(
global_step,
actor,
qnet,
qnet_target,
obs_normalizer,
critic_obs_normalizer,
args,
f"models/{run_name}_{global_step}.pt",
)
global_step += 1
pbar.update(1)
save_params(
global_step,
actor,
qnet,
qnet_target,
obs_normalizer,
critic_obs_normalizer,
args,
f"models/{run_name}_final.pt",
)
if __name__ == "__main__":
main()
+543
View File
@@ -0,0 +1,543 @@
import os
from dataclasses import dataclass
import tyro
@dataclass
class BaseArgs:
# Default hyperparameters -- specifically for HumanoidBench
# See MuJoCoPlaygroundArgs for default hyperparameters for MuJoCo Playground
# See IsaacLabArgs for default hyperparameters for IsaacLab
env_name: str = "HumanoidRun"
"""the id of the environment"""
agent: str = "fasttd3"
"""the agent to use: currently support [fasttd3, fasttd3_simbav2]"""
seed: int = 1
"""seed of the experiment"""
torch_deterministic: bool = True
"""if toggled, `torch.backends.cudnn.deterministic=False`"""
cuda: bool = True
"""if toggled, cuda will be enabled by default"""
device_rank: int = 0
"""the rank of the device"""
exp_name: str = os.path.basename(__file__)[: -len(".py")]
"""the name of this experiment"""
project: str = "FastTD3"
"""the project name"""
use_wandb: bool = True
"""whether to use wandb"""
checkpoint_path: str = None
"""the path to the checkpoint file"""
num_envs: int = 128
"""the number of environments to run in parallel"""
num_eval_envs: int = 128
"""the number of evaluation environments to run in parallel (only valid for MuJoCo Playground)"""
total_timesteps: int = 50000
"""total timesteps of the experiments"""
critic_learning_rate: float = 3e-4
"""the learning rate of the critic"""
actor_learning_rate: float = 3e-4
"""the learning rate for the actor"""
critic_learning_rate_end: float = 3e-4
"""the learning rate of the critic at the end of training"""
actor_learning_rate_end: float = 3e-4
"""the learning rate for the actor at the end of training"""
buffer_size: int = 1024 * 50
"""the replay memory buffer size"""
num_steps: int = 1
"""the number of steps to use for the multi-step return"""
gamma: float = 0.99
"""the discount factor gamma"""
tau: float = 0.1
"""target smoothing coefficient (default: 0.005)"""
batch_size: int = 32768
"""the batch size of sample from the replay memory"""
policy_noise: float = 0.001
"""the scale of policy noise"""
std_min: float = 0.001
"""the minimum scale of noise"""
std_max: float = 0.4
"""the maximum scale of noise"""
learning_starts: int = 10
"""timestep to start learning"""
policy_frequency: int = 2
"""the frequency of training policy (delayed)"""
noise_clip: float = 0.5
"""noise clip parameter of the Target Policy Smoothing Regularization"""
num_updates: int = 2
"""the number of updates to perform per step"""
init_scale: float = 0.01
"""the scale of the initial parameters"""
num_atoms: int = 101
"""the number of atoms"""
v_min: float = -250.0
"""the minimum value of the support"""
v_max: float = 250.0
"""the maximum value of the support"""
critic_hidden_dim: int = 1024
"""the hidden dimension of the critic network"""
actor_hidden_dim: int = 512
"""the hidden dimension of the actor network"""
critic_num_blocks: int = 2
"""(SimbaV2 only) the number of blocks in the critic network"""
actor_num_blocks: int = 1
"""(SimbaV2 only) the number of blocks in the actor network"""
use_cdq: bool = True
"""whether to use Clipped Double Q-learning"""
measure_burnin: int = 3
"""Number of burn-in iterations for speed measure."""
eval_interval: int = 2500
"""the interval to evaluate the model"""
render_interval: int = 500000
"""the interval to render the model"""
compile: bool = True
"""whether to use torch.compile."""
compile_mode: str = "reduce-overhead"
"""the mode of torch.compile."""
obs_normalization: bool = True
"""whether to enable observation normalization"""
reward_normalization: bool = False
"""whether to enable reward normalization"""
use_grad_norm_clipping: bool = False
"""whether to use gradient norm clipping."""
max_grad_norm: float = 0.0
"""the maximum gradient norm"""
amp: bool = True
"""whether to use amp"""
amp_dtype: str = "bf16"
"""the dtype of the amp"""
disable_bootstrap: bool = False
"""Whether to disable bootstrap in the critic learning"""
use_domain_randomization: bool = False
"""(Playground only) whether to use domain randomization"""
use_push_randomization: bool = False
"""(Playground only) whether to use push randomization"""
use_tuned_reward: bool = False
"""(Playground only) Use tuned reward for G1"""
action_bounds: float = 1.0
"""(IsaacLab only) the bounds of the action space (-action_bounds, action_bounds)"""
task_embedding_dim: int = 32
"""the dimension of the task embedding"""
weight_decay: float = 0.1
"""the weight decay of the optimizer"""
save_interval: int = 5000
"""the interval to save the model"""
def get_args():
"""
Parse command-line arguments and return the appropriate Args instance based on env_name.
"""
# First, parse all arguments using the base Args class
base_args = tyro.cli(BaseArgs)
# Map environment names to their specific Args classes
# For tasks not here, default hyperparameters are used
# See below links for available task list
# - HumanoidBench (https://arxiv.org/abs/2403.10506)
# - IsaacLab (https://isaac-sim.github.io/IsaacLab/main/source/overview/environments.html)
# - MuJoCo Playground (https://arxiv.org/abs/2502.08844)
env_to_args_class = {
# HumanoidBench
# NOTE: These tasks are not full list of HumanoidBench tasks
"h1hand-reach-v0": H1HandReachArgs,
"h1hand-balance-simple-v0": H1HandBalanceSimpleArgs,
"h1hand-balance-hard-v0": H1HandBalanceHardArgs,
"h1hand-pole-v0": H1HandPoleArgs,
"h1hand-truck-v0": H1HandTruckArgs,
"h1hand-maze-v0": H1HandMazeArgs,
"h1hand-push-v0": H1HandPushArgs,
"h1hand-basketball-v0": H1HandBasketballArgs,
"h1hand-window-v0": H1HandWindowArgs,
"h1hand-package-v0": H1HandPackageArgs,
# MuJoCo Playground
# NOTE: These tasks are not full list of MuJoCo Playground tasks
"G1JoystickFlatTerrain": G1JoystickFlatTerrainArgs,
"G1JoystickRoughTerrain": G1JoystickRoughTerrainArgs,
"T1JoystickFlatTerrain": T1JoystickFlatTerrainArgs,
"T1JoystickRoughTerrain": T1JoystickRoughTerrainArgs,
"LeapCubeReorient": LeapCubeReorientArgs,
"LeapCubeRotateZAxis": LeapCubeRotateZAxisArgs,
"Go1JoystickFlatTerrain": Go1JoystickFlatTerrainArgs,
"Go1JoystickRoughTerrain": Go1JoystickRoughTerrainArgs,
"Go1Getup": Go1GetupArgs,
"CheetahRun": CheetahRunArgs, # NOTE: Example config for DeepMind Control Suite
# IsaacLab
# NOTE: These tasks are not full list of IsaacLab tasks
"Isaac-Lift-Cube-Franka-v0": IsaacLiftCubeFrankaArgs,
"Isaac-Open-Drawer-Franka-v0": IsaacOpenDrawerFrankaArgs,
"Isaac-Velocity-Flat-H1-v0": IsaacVelocityFlatH1Args,
"Isaac-Velocity-Flat-G1-v0": IsaacVelocityFlatG1Args,
"Isaac-Velocity-Rough-H1-v0": IsaacVelocityRoughH1Args,
"Isaac-Velocity-Rough-G1-v0": IsaacVelocityRoughG1Args,
"Isaac-Repose-Cube-Allegro-Direct-v0": IsaacReposeCubeAllegroDirectArgs,
"Isaac-Repose-Cube-Shadow-Direct-v0": IsaacReposeCubeShadowDirectArgs,
# MTBench
"MTBench-meta-world-v2-mt10": MetaWorldMT10Args,
"MTBench-meta-world-v2-mt50": MetaWorldMT50Args,
}
# If the provided env_name has a specific Args class, use it
if base_args.env_name in env_to_args_class:
specific_args_class = env_to_args_class[base_args.env_name]
# Re-parse with the specific class, maintaining any user overrides
specific_args = tyro.cli(specific_args_class)
return specific_args
if base_args.env_name.startswith("h1hand-") or base_args.env_name.startswith("h1-"):
# HumanoidBench
specific_args = tyro.cli(HumanoidBenchArgs)
elif base_args.env_name.startswith("Isaac-"):
# IsaacLab
specific_args = tyro.cli(IsaacLabArgs)
elif base_args.env_name.startswith("MTBench-"):
# MTBench
specific_args = tyro.cli(MTBenchArgs)
else:
# MuJoCo Playground
specific_args = tyro.cli(MuJoCoPlaygroundArgs)
return specific_args
@dataclass
class HumanoidBenchArgs(BaseArgs):
# See HumanoidBench (https://arxiv.org/abs/2403.10506) for available task list
total_timesteps: int = 100000
@dataclass
class H1HandReachArgs(HumanoidBenchArgs):
env_name: str = "h1hand-reach-v0"
v_min: float = -2000.0
v_max: float = 2000.0
@dataclass
class H1HandBalanceSimpleArgs(HumanoidBenchArgs):
env_name: str = "h1hand-balance-simple-v0"
total_timesteps: int = 200000
@dataclass
class H1HandBalanceHardArgs(HumanoidBenchArgs):
env_name: str = "h1hand-balance-hard-v0"
total_timesteps: int = 1000000
@dataclass
class H1HandPoleArgs(HumanoidBenchArgs):
env_name: str = "h1hand-pole-v0"
total_timesteps: int = 150000
@dataclass
class H1HandTruckArgs(HumanoidBenchArgs):
env_name: str = "h1hand-truck-v0"
total_timesteps: int = 500000
@dataclass
class H1HandMazeArgs(HumanoidBenchArgs):
env_name: str = "h1hand-maze-v0"
v_min: float = -1000.0
v_max: float = 1000.0
@dataclass
class H1HandPushArgs(HumanoidBenchArgs):
env_name: str = "h1hand-push-v0"
v_min: float = -1000.0
v_max: float = 1000.0
total_timesteps: int = 1000000
@dataclass
class H1HandBasketballArgs(HumanoidBenchArgs):
env_name: str = "h1hand-basketball-v0"
v_min: float = -2000.0
v_max: float = 2000.0
total_timesteps: int = 250000
@dataclass
class H1HandWindowArgs(HumanoidBenchArgs):
env_name: str = "h1hand-window-v0"
total_timesteps: int = 250000
@dataclass
class H1HandPackageArgs(HumanoidBenchArgs):
env_name: str = "h1hand-package-v0"
v_min: float = -10000.0
v_max: float = 10000.0
@dataclass
class MuJoCoPlaygroundArgs(BaseArgs):
# Default hyperparameters for many of Playground environments
v_min: float = -150.0
v_max: float = 150.0
buffer_size: int = 1024 * 10
num_envs: int = 1024
num_eval_envs: int = 1024
gamma: float = 0.99
@dataclass
class MTBenchArgs(BaseArgs):
# Default hyperparameters for MTBench
reward_normalization: bool = True
v_min: float = -10.0
v_max: float = 10.0
buffer_size: int = 2048 # 2K is usually enough for MTBench
num_envs: int = 4096
num_eval_envs: int = 4096
gamma: float = 0.97
num_steps: int = 8
@dataclass
class MetaWorldMT10Args(MTBenchArgs):
# This config achieves 97 ~ 98% success rate within 10k steps (15-20 mins on A100)
env_name: str = "MTBench-meta-world-v2-mt10"
num_envs: int = 4096
num_eval_envs: int = 4096
num_steps: int = 8
gamma: float = 0.97
@dataclass
class MetaWorldMT50Args(MTBenchArgs):
# FastTD3 + SimbaV2 achieves >90% success rate within 20k steps (80 mins on A100)
# Performance further improves with more training steps, slowly.
env_name: str = "MTBench-meta-world-v2-mt50"
num_envs: int = 8192
num_eval_envs: int = 8192
num_steps: int = 8
gamma: float = 0.99
@dataclass
class G1JoystickFlatTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "G1JoystickFlatTerrain"
total_timesteps: int = 100000
v_min: float = -10.0
v_max: float = 10.0
buffer_size: int = 128 # 1024 * 10
num_envs: int = 1024
num_eval_envs: int = 1024
gamma: float = 0.97
critic_hidden_dim: int = 1024
batch_size: int = 8129
@dataclass
class G1JoystickRoughTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "G1JoystickRoughTerrain"
total_timesteps: int = 100000
v_min: float = -10.0
v_max: float = 10.0
buffer_size: int = 128 # 1024 * 10
num_envs: int = 1024
num_eval_envs: int = 1024
gamma: float = 0.97
critic_hidden_dim: int = 1024
batch_size: int = 8129
@dataclass
class T1JoystickFlatTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "T1JoystickFlatTerrain"
total_timesteps: int = 100000
v_min: float = -10.0
v_max: float = 10.0
buffer_size: int = 128 # 1024 * 10
num_envs: int = 1024
num_eval_envs: int = 1024
gamma: float = 0.97
critic_hidden_dim: int = 1024
batch_size: int = 8129
@dataclass
class T1JoystickRoughTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "T1JoystickRoughTerrain"
total_timesteps: int = 100000
v_min: float = -10.0
v_max: float = 10.0
buffer_size: int = 128 # 1024 * 10
num_envs: int = 1024
num_eval_envs: int = 1024
gamma: float = 0.97
critic_hidden_dim: int = 1024
batch_size: int = 8129
@dataclass
class T1LowDofJoystickFlatTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "T1LowDofJoystickFlatTerrain"
total_timesteps: int = 1000000
@dataclass
class T1LowDofJoystickRoughTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "T1LowDofJoystickRoughTerrain"
total_timesteps: int = 1000000
@dataclass
class CheetahRunArgs(MuJoCoPlaygroundArgs):
# NOTE: This config will work for most DMC tasks, though we haven't tested DMC extensively.
# Future research can consider using LayerNorm as we find it sometimes works better for DMC tasks.
env_name: str = "CheetahRun"
num_steps: int = 3
v_min: float = -500.0
v_max: float = 500.0
std_min: float = 0.1
policy_noise: float = 0.1
@dataclass
class Go1JoystickFlatTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "Go1JoystickFlatTerrain"
total_timesteps: int = 50000
std_min: float = 0.2
std_max: float = 0.8
policy_noise: float = 0.2
num_updates: int = 8
@dataclass
class Go1JoystickRoughTerrainArgs(MuJoCoPlaygroundArgs):
env_name: str = "Go1JoystickRoughTerrain"
total_timesteps: int = 50000
std_min: float = 0.2
std_max: float = 0.8
policy_noise: float = 0.2
num_updates: int = 8
@dataclass
class Go1GetupArgs(MuJoCoPlaygroundArgs):
env_name: str = "Go1Getup"
total_timesteps: int = 50000
std_min: float = 0.2
std_max: float = 0.8
policy_noise: float = 0.2
num_updates: int = 8
@dataclass
class LeapCubeReorientArgs(MuJoCoPlaygroundArgs):
env_name: str = "LeapCubeReorient"
num_steps: int = 3
gamma: float = 0.99
policy_noise: float = 0.2
v_min: float = -50.0
v_max: float = 50.0
use_cdq: bool = False
@dataclass
class LeapCubeRotateZAxisArgs(MuJoCoPlaygroundArgs):
env_name: str = "LeapCubeRotateZAxis"
num_steps: int = 1
policy_noise: float = 0.2
gamma: float = 0.99
v_min: float = -10.0
v_max: float = 10.0
use_cdq: bool = False
@dataclass
class IsaacLabArgs(BaseArgs):
v_min: float = -10.0
v_max: float = 10.0
buffer_size: int = 1024 * 10
num_envs: int = 4096
num_eval_envs: int = 4096
action_bounds: float = 1.0
std_max: float = 0.4
num_atoms: int = 251
render_interval: int = 0 # IsaacLab does not support rendering in our codebase
total_timesteps: int = 100000
@dataclass
class IsaacLiftCubeFrankaArgs(IsaacLabArgs):
# Value learning is unstable for Lift Cube task Due to brittle reward shaping
# Therefore, we need to disable bootstrap from 'reset_obs' in IsaacLab
# Higher UTD works better for manipulation tasks
env_name: str = "Isaac-Lift-Cube-Franka-v0"
num_updates: int = 8
v_min: float = -50.0
v_max: float = 50.0
std_max: float = 0.8
num_envs: int = 1024
num_eval_envs: int = 1024
action_bounds: float = 3.0
disable_bootstrap: bool = True
total_timesteps: int = 20000
@dataclass
class IsaacOpenDrawerFrankaArgs(IsaacLabArgs):
# Higher UTD works better for manipulation tasks
env_name: str = "Isaac-Open-Drawer-Franka-v0"
v_min: float = -50.0
v_max: float = 50.0
num_updates: int = 8
action_bounds: float = 3.0
total_timesteps: int = 20000
@dataclass
class IsaacVelocityFlatH1Args(IsaacLabArgs):
env_name: str = "Isaac-Velocity-Flat-H1-v0"
num_steps: int = 8
num_updates: int = 4
total_timesteps: int = 75000
@dataclass
class IsaacVelocityFlatG1Args(IsaacLabArgs):
env_name: str = "Isaac-Velocity-Flat-G1-v0"
num_steps: int = 8
num_updates: int = 4
total_timesteps: int = 50000
@dataclass
class IsaacVelocityRoughH1Args(IsaacLabArgs):
env_name: str = "Isaac-Velocity-Rough-H1-v0"
num_steps: int = 8
num_updates: int = 4
buffer_size: int = 1024 * 5 # To reduce memory usage
total_timesteps: int = 50000
@dataclass
class IsaacVelocityRoughG1Args(IsaacLabArgs):
env_name: str = "Isaac-Velocity-Rough-G1-v0"
num_steps: int = 8
num_updates: int = 4
buffer_size: int = 1024 * 5 # To reduce memory usage
total_timesteps: int = 50000
@dataclass
class IsaacReposeCubeAllegroDirectArgs(IsaacLabArgs):
env_name: str = "Isaac-Repose-Cube-Allegro-Direct-v0"
total_timesteps: int = 100000
v_min: float = -500.0
v_max: float = 500.0
@dataclass
class IsaacReposeCubeShadowDirectArgs(IsaacLabArgs):
env_name: str = "Isaac-Repose-Cube-Shadow-Direct-v0"
total_timesteps: int = 100000
v_min: float = -500.0
v_max: float = 500.0
+735
View File
@@ -0,0 +1,735 @@
from dataclasses import dataclass, replace
import functools
import os
import random
import sys
import copy
import time
import numpy as np
import tqdm
from omegaconf import DictConfig, OmegaConf
import wandb
from reppo_alg.torchrl.reppo import EmpiricalNormalization, hl_gauss
try:
# Required for avoiding IsaacGym import error
import isaacgym
except ImportError:
pass
import hydra
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torchinfo import summary
from tensordict import TensorDict
from torch.amp import GradScaler
from reppo_alg.torchrl.envs import make_envs
from reppo_alg.network_utils.torch_models import Actor, Critic
torch.set_float32_matmul_precision("medium")
os.environ["TORCHDYNAMO_INLINE_INBUILT_NN_MODULES"] = "1"
os.environ["OMP_NUM_THREADS"] = "1"
if sys.platform != "darwin":
os.environ["MUJOCO_GL"] = "egl"
else:
os.environ["MUJOCO_GL"] = "glfw"
os.environ["XLA_PYTHON_CLIENT_PREALLOCATE"] = "false"
os.environ["JAX_DEFAULT_MATMUL_PRECISION"] = "highest"
@dataclass(slots=True)
class TrainState:
device: torch.device
obs: torch.Tensor
critic_obs: torch.Tensor
actor: Actor
old_actor: Actor
critic: Critic
normalizer: EmpiricalNormalization
critic_normalizer: EmpiricalNormalization
actor_optimizer: optim.Optimizer
critic_optimizer: optim.Optimizer
scaler: GradScaler
def compile(self):
self.actor.compile()
self.old_actor.compile()
self.critic.compile()
self.normalizer.compile()
self.critic_normalizer.compile()
def get_autocast_context(cfg: DictConfig):
amp_enabled = (
cfg.platform.amp_enabled and cfg.platform.cuda and torch.cuda.is_available()
)
amp_device = (
"cuda"
if cfg.platform.cuda and torch.cuda.is_available()
else "mps"
if cfg.platform.cuda and torch.backends.mps.is_available()
else "cpu"
)
amp_dtype = torch.bfloat16 if cfg.platform.amp_dtype == "bf16" else torch.float32
return functools.partial(
torch.amp.autocast,
device_type=amp_device,
dtype=amp_dtype,
enabled=amp_enabled,
)
def make_collect_fn(cfg: DictConfig, env):
autocast = get_autocast_context(cfg)
asymmetric_obs = env.asymmetric_obs
def collect_fn(
train_state: TrainState,
) -> tuple[TrainState, TensorDict, list[dict]]:
transitions = []
info_list = []
obs = train_state.obs
critic_obs = train_state.critic_obs
for _ in range(cfg.hyperparameters.num_steps):
with autocast():
norm_obs = train_state.normalizer(obs)
norm_critic_obs = train_state.critic_normalizer(critic_obs)
with torch.inference_mode():
pi, _, _, _ = train_state.actor(norm_obs)
actions = pi.sample()
next_obs, rewards, dones, truncations, infos = env.step(actions)
if asymmetric_obs:
next_critic_obs = infos["observations"]["critic"]
else:
next_critic_obs = next_obs
with (
torch.inference_mode(),
autocast(),
):
if (
cfg.env.get("has_final_obs", False)
and cfg.env.get("partial_reset", False)
and "final_observation" in infos
):
_next_obs = infos["final_observation"]
_next_critic_obs = _next_obs
else:
_next_obs = next_obs
_next_critic_obs = next_critic_obs
norm_next_obs = train_state.normalizer(_next_obs)
next_pi, _, temperature, _ = train_state.actor(norm_next_obs)
next_actions = next_pi.sample()
next_log_probs = next_pi.log_prob(
next_actions.clip(-1 + 1e-6, 1 - 1e-6)
).sum(-1)
norm_next_critic_obs = train_state.critic_normalizer(_next_critic_obs)
next_value, _, _, next_embedding = train_state.critic(
norm_next_critic_obs, next_actions
)
rewards = (
rewards - cfg.hyperparameters.gamma * next_log_probs * temperature
)
transitions.append(
TensorDict(
{
"observations": norm_obs,
"critic_observations": norm_critic_obs,
"actions": actions,
"log_probs": pi.log_prob(actions.clip(-0.999, 0.999)).sum(-1),
"rewards": rewards.unsqueeze(-1),
"next_embeddings": next_embedding,
"next_values": next_value.unsqueeze(-1),
"dones": dones.unsqueeze(-1).float(),
"truncations": truncations.unsqueeze(-1).float(),
},
batch_size=(env.num_envs,),
)
)
info_list.append(infos)
obs = next_obs
critic_obs = next_critic_obs
train_state = replace(train_state, obs=obs, critic_obs=critic_obs)
return (
train_state,
torch.stack(transitions, dim=0),
info_list,
)
return collect_fn
def make_postprocess_fn(cfg: DictConfig, env):
@torch.compiler.disable()
def compute_gve(rewards, dones, truncated, next_values, device: torch.device):
gves = []
last_gve = 0
truncated[-1] = 1.0
for t in reversed(range(cfg.hyperparameters.num_steps)):
lambda_sum = (
cfg.hyperparameters.lmbda * last_gve
+ (1.0 - cfg.hyperparameters.lmbda) * next_values[t]
)
delta = cfg.hyperparameters.gamma * torch.where(
truncated[t].bool(), next_values[t], (1.0 - dones[t]) * lambda_sum
)
last_gve = rewards[t] + delta
gves.insert(0, last_gve)
return gves
def postprocess(train_state: TrainState, transition: TensorDict):
gve = compute_gve(
rewards=transition["rewards"],
dones=transition["dones"],
truncated=transition["truncations"],
next_values=transition["next_values"],
device=train_state.device,
)
# Flatten all time and environment dimensions into a single batch dimension
data = TensorDict(
{
"observations": transition["observations"],
"critic_observations": transition["critic_observations"],
"actions": transition["actions"],
"rewards": transition["rewards"],
"next_embeddings": transition["next_embeddings"],
"next_values": transition["next_values"],
"dones": transition["dones"],
"truncations": transition["truncations"],
"gve": torch.stack(gve),
},
batch_size=(
cfg.hyperparameters.num_steps,
cfg.hyperparameters.num_envs,
),
device=train_state.device,
)
return data.float().flatten(0, 1).detach()
return postprocess
def make_critic_update_fn(cfg: DictConfig, train_state: TrainState):
autocast = get_autocast_context(cfg)
def update(data: TensorDict):
qnet = train_state.critic
q_optimizer = train_state.critic_optimizer
with autocast():
critic_observations = data["critic_observations"]
actions = data["actions"]
targets = data["gve"]
target_embeddings = data["next_embeddings"]
truncations = data["truncations"].squeeze(-1)
if cfg.env.get("partial_reset", False):
truncation_mask = torch.ones_like(
truncations, dtype=torch.bool, device=train_state.device
)
else:
truncation_mask = 1.0 - truncations
qf_target_dist = hl_gauss(
targets,
cfg.hyperparameters.vmin,
cfg.hyperparameters.vmax,
cfg.hyperparameters.num_bins,
)
_, qf1, embedding, _ = qnet(critic_observations, actions)
qf_loss = -(
truncation_mask
* torch.sum(qf_target_dist * F.log_softmax(qf1, dim=-1), dim=-1)
).mean()
embedding_loss = (
truncation_mask
* F.mse_loss(
embedding,
target_embeddings,
reduction="none",
).mean(dim=-1)
).mean()
qf_loss = qf_loss + cfg.hyperparameters.aux_loss_mult * embedding_loss
q_optimizer.zero_grad(set_to_none=True)
train_state.scaler.scale(qf_loss).backward()
train_state.scaler.unscale_(q_optimizer)
critic_grad_norm = torch.nn.utils.clip_grad_norm_(
qnet.parameters(), max_norm=cfg.hyperparameters.max_grad_norm
)
train_state.scaler.step(q_optimizer)
train_state.scaler.update()
logs_dict = {
"critic_grad_norm": critic_grad_norm.detach(),
"qf_loss": qf_loss.detach(),
"qf_max": targets.max().detach(),
"qf_min": targets.min().detach(),
"qf_mean": targets.mean().detach(),
"embedding_loss": embedding_loss.detach(),
}
return logs_dict
return update
def make_actor_update_fn(cfg: DictConfig, train_state: TrainState):
autocast = get_autocast_context(cfg)
def update(data: TensorDict):
actor = train_state.actor
old_actor = train_state.old_actor
qnet = train_state.critic
actor_optimizer = train_state.actor_optimizer
scaler = train_state.scaler
critic_obs = data["critic_observations"]
with autocast():
pi, _, temperature, beta = actor(data["observations"])
actions = pi.rsample()
log_probs = pi.log_prob(actions.clip(-1 + 1e-6, 1 - 1e-6)).sum(-1)
entropy = -log_probs
qf, _, _, _ = qnet(critic_obs, actions)
actor_loss = -qf + temperature.detach() * log_probs
# compute KL
old_pi, _, _, _ = old_actor(data["observations"])
old_pi_actions = old_pi.sample((16,)).clip(-1 + 1e-6, 1 - 1e-6)
old_log_probs = old_pi.log_prob(old_pi_actions).sum(-1).mean(0)
new_pi_log_probs = pi.log_prob(old_pi_actions).sum(-1).mean(0)
kl = old_log_probs - new_pi_log_probs
if cfg.hyperparameters.actor_kl_clip_mode == "clipped":
actor_loss = torch.where(
kl < cfg.hyperparameters.kl_bound,
actor_loss,
kl * beta.detach(),
).mean()
elif cfg.hyperparameters.actor_kl_clip_mode == "full":
actor_loss = actor_loss + kl * beta.detach()
elif cfg.hyperparameters.actor_kl_clip_mode == "value":
actor_loss = actor_loss
else:
raise ValueError(
f"Unknown actor kl clip mode: {cfg.hyperparameters.actor_kl_clip_mode}"
)
# temperature updates
target_entropy = (
actions.shape[-1] * cfg.hyperparameters.ent_target_mult
) # -0.5 * np.prod(envs.action_space.shape)
entropy_loss = (target_entropy + entropy).detach().mean() * temperature
lagrangian_loss = (
-beta * (kl - cfg.hyperparameters.kl_bound).mean().detach()
)
actor_loss = (actor_loss + entropy_loss + lagrangian_loss).mean()
actor_optimizer.zero_grad(set_to_none=True)
scaler.scale(actor_loss).backward()
scaler.unscale_(actor_optimizer)
actor_grad_norm = torch.nn.utils.clip_grad_norm_(
actor.parameters(), max_norm=cfg.hyperparameters.max_grad_norm
)
scaler.step(actor_optimizer)
scaler.update()
logs_dict = {
"actor_grad_norm": actor_grad_norm.detach(),
"actor_loss": actor_loss.detach(),
"kl": kl.detach(),
"entropy": entropy.detach(),
"temperature": temperature.detach(),
"lagrangian": beta.detach(),
"entropy_loss": entropy_loss.detach(),
"lagrangian_loss": lagrangian_loss.detach(),
}
return logs_dict
return update
def make_evaluate_fn(cfg: DictConfig, eval_envs):
autocast = get_autocast_context(cfg)
@torch.inference_mode()
def evaluate(
train_state: TrainState, stochastic_eval: bool = False
) -> tuple[int | float | bool, int | float | bool]:
train_state.normalizer.eval()
num_eval_envs = eval_envs.num_envs
episode_returns = torch.zeros(num_eval_envs, device=train_state.device)
episode_lengths = torch.zeros(num_eval_envs, device=train_state.device)
done_masks = torch.zeros(
num_eval_envs, dtype=torch.bool, device=train_state.device
)
if cfg.env.type == "isaaclab" or cfg.env.asymmetric_observation:
obs, _ = eval_envs.reset(random_start_init=False)
else:
obs = eval_envs.reset()
# Run for a fixed number of steps
for i in range(eval_envs.max_episode_steps):
with autocast():
obs = train_state.normalizer(obs)
action_dist, det_actions, _, _ = train_state.actor(obs)
if stochastic_eval:
actions = action_dist.sample()
else:
actions = det_actions
next_obs, rewards, dones, _, infos = eval_envs.step(actions)
episode_returns = torch.where(
~done_masks, episode_returns + rewards, episode_returns
)
episode_lengths = torch.where(
~done_masks, episode_lengths + 1, episode_lengths
)
done_masks = torch.logical_or(done_masks, dones)
if done_masks.all():
break
obs = next_obs
train_state.normalizer.train()
if cfg.env.type == "maniskill":
# combine log_infos
info = {
"info_return": infos["log_info"]["return"].mean(),
"episode_len": infos["log_info"]["episode_len"].float().mean(),
"success": infos["log_info"]["success"].float().mean(),
"return": episode_returns.mean().item(),
}
else:
info = {}
return episode_returns.mean().item(), episode_lengths.mean().item(), info
return evaluate
def configure_platform(cfg: DictConfig) -> DictConfig:
cfg.platform.amp_enabled = (
cfg.platform.amp_enabled and cfg.platform.cuda and torch.cuda.is_available()
)
cfg.platform.amp_device = (
"cuda"
if cfg.platform.cuda and torch.cuda.is_available()
else "mps"
if cfg.platform.cuda and torch.backends.mps.is_available()
else "cpu"
)
return cfg
@hydra.main(
version_base=None,
config_path="../../config",
config_name="reppo",
)
def main(cfg):
cfg = configure_platform(cfg)
run_name = f"{cfg.env.name}_torch_{cfg.seed}"
scaler = GradScaler(
enabled=cfg.platform.amp_enabled and cfg.platform.amp_dtype == torch.float16
)
num_batches = cfg.hyperparameters.num_mini_batches
batch_size = (
cfg.hyperparameters.num_envs * cfg.hyperparameters.num_steps // num_batches
)
wandb.init(
project=cfg.wandb.project,
name=run_name,
config=OmegaConf.to_container(cfg),
save_code=True,
)
random.seed(cfg.seed)
np.random.seed(cfg.seed)
torch.manual_seed(cfg.seed)
torch.backends.cudnn.deterministic = cfg.platform.torch_deterministic
if not cfg.platform.cuda:
device = torch.device("cpu")
else:
if torch.cuda.is_available():
device = torch.device(f"cuda:{cfg.platform.device_rank}")
elif torch.backends.mps.is_available():
device = torch.device(f"mps:{cfg.platform.device_rank}")
else:
raise ValueError("No GPU available")
print(f"Using device: {device}")
envs, eval_envs = make_envs(cfg=cfg, device=device, seed=cfg.seed)
n_act = envs.num_actions
n_obs = envs.num_obs if isinstance(envs.num_obs, int) else envs.num_obs[0]
if envs.asymmetric_obs:
n_critic_obs = (
envs.num_privileged_obs
if isinstance(envs.num_privileged_obs, int)
else envs.num_privileged_obs[0]
)
else:
n_critic_obs = n_obs
if cfg.hyperparameters.normalize_env:
obs_normalizer = EmpiricalNormalization(shape=n_obs, device=device)
critic_obs_normalizer = EmpiricalNormalization(
shape=n_critic_obs, device=device
)
else:
obs_normalizer = nn.Identity()
critic_obs_normalizer = nn.Identity()
actor = Actor(
n_obs=n_obs,
n_act=n_act,
ent_start=cfg.hyperparameters.ent_start,
kl_start=cfg.hyperparameters.kl_start,
hidden_dim=cfg.hyperparameters.actor_hidden_dim,
use_norm=cfg.hyperparameters.use_actor_norm,
layers=cfg.hyperparameters.num_actor_layers,
min_std=cfg.hyperparameters.actor_min_std,
device=device,
)
old_actor = copy.deepcopy(actor)
qnet = Critic(
n_obs=n_critic_obs,
n_act=n_act,
num_atoms=cfg.hyperparameters.num_bins,
vmin=cfg.hyperparameters.vmin,
vmax=cfg.hyperparameters.vmax,
hidden_dim=cfg.hyperparameters.critic_hidden_dim,
use_norm=cfg.hyperparameters.use_critic_norm,
use_encoder_norm=False,
encoder_layers=cfg.hyperparameters.num_critic_encoder_layers,
head_layers=cfg.hyperparameters.num_critic_head_layers,
pred_layers=cfg.hyperparameters.num_critic_pred_layers,
device=device,
)
q_optimizer = optim.AdamW(
list(qnet.parameters()),
lr=torch.tensor(cfg.hyperparameters.lr, device=device),
)
actor_optimizer = optim.AdamW(
list(actor.parameters()),
lr=torch.tensor(cfg.hyperparameters.lr, device=device),
)
if envs.asymmetric_obs:
obs, critic_obs = envs.reset_with_critic_obs()
critic_obs = torch.as_tensor(critic_obs, device=device, dtype=torch.float)
else:
obs = envs.reset()
critic_obs = obs
train_state = TrainState(
obs=obs,
critic_obs=critic_obs,
actor=actor,
old_actor=old_actor,
critic=qnet,
normalizer=obs_normalizer,
critic_normalizer=critic_obs_normalizer,
actor_optimizer=actor_optimizer,
critic_optimizer=q_optimizer,
device=device,
scaler=scaler,
)
print(
summary(
train_state.critic,
input_data=(critic_obs[:1], torch.zeros((1, n_act), device=device)),
depth=10,
)
)
print(summary(train_state.actor, input_data=(obs[:1],), depth=10))
# create functions
collect_fn = make_collect_fn(cfg, envs)
postprocess_fn = make_postprocess_fn(cfg, envs)
update_critic = make_critic_update_fn(cfg, train_state)
update_actor = make_actor_update_fn(cfg, train_state)
evaluate = make_evaluate_fn(cfg, eval_envs)
if cfg.platform.compile:
mode = "max-autotune-no-cudagraphs"
update_critic = torch.compile(update_critic, mode=mode)
update_actor = torch.compile(update_actor, mode=mode)
postprocess_fn = torch.compile(postprocess_fn, mode=mode)
train_state.compile()
# TODO: Support checkpoint loading
# if cfg.checkpoint_path:
# # Load checkpoint if specified
# torch_checkpoint = torch.load(
# f"{cfg.checkpoint_path}", map_location=device, weights_only=False
# )
# actor.load_state_dict(torch_checkpoint["actor_state_dict"])
# obs_normalizer.load_state_dict(torch_checkpoint["obs_normalizer_state"])
# critic_obs_normalizer.load_state_dict(
# torch_checkpoint["critic_obs_normalizer_state"]
# )
# qnet.load_state_dict(torch_checkpoint["qnet_state_dict"])
# qnet_target.load_state_dict(torch_checkpoint["qnet_target_state_dict"])
# global_step = torch_checkpoint["global_step"]
# else:
global_step = 0
total_env_steps = (
cfg.hyperparameters.total_time_steps
// (cfg.hyperparameters.num_envs * cfg.hyperparameters.num_steps)
+ 1
)
pbar = tqdm.tqdm(total=cfg.hyperparameters.total_time_steps, initial=global_step)
start_time = None
desc = ""
eval_interval = total_env_steps // cfg.hyperparameters.num_eval
stochastic_eval = cfg.env.get("stochastic_eval", False)
while global_step < total_env_steps:
if start_time is None and global_step >= cfg.measure_burnin:
start_time = time.time()
measure_burnin = global_step
train_state, transition, infos = collect_fn(train_state)
data = postprocess_fn(train_state, transition)
for _ in range(cfg.hyperparameters.num_epochs):
indices = torch.randperm(
cfg.hyperparameters.num_envs * cfg.hyperparameters.num_steps,
device=device,
)
data = data[indices].contiguous()
for j in range(num_batches):
mini_batch = data[j * batch_size : (j + 1) * batch_size]
critic_logs_dict = update_critic(mini_batch)
actor_logs_dict = update_actor(mini_batch)
logs_dict = {
**critic_logs_dict,
**actor_logs_dict,
}
for param, target_param in zip(actor.parameters(), old_actor.parameters()):
target_param.data.copy_(param.data)
if start_time is not None:
# @TODO: shouldn't that be env_steps per second?
speed = (
cfg.hyperparameters.num_envs
* cfg.hyperparameters.num_steps
* (global_step - measure_burnin)
/ (time.time() - start_time)
)
pbar.set_description(f"{speed: 4.4f} sps, " + desc)
with torch.inference_mode():
logs = {
"critic/qf_loss": logs_dict["qf_loss"].mean(),
"critic/qf_max": logs_dict["qf_max"].mean(),
"critic/qf_min": logs_dict["qf_min"].mean(),
"critic/qf_mean": logs_dict["qf_mean"].mean(),
"critic/embedding_loss": logs_dict["embedding_loss"].mean(),
"critic/critic_grad_norm": logs_dict["critic_grad_norm"].mean(),
"actor/actor_loss": logs_dict["actor_loss"].mean(),
"actor/actor_grad_norm": logs_dict["actor_grad_norm"].mean(),
"actor/kl": logs_dict["kl"].mean(),
"actor/entropy": logs_dict["entropy"].mean(),
"actor/temperature": logs_dict["temperature"].mean(),
"actor/lagrangian": logs_dict["lagrangian"].mean(),
"actor/entropy_loss": logs_dict["entropy_loss"].mean(),
"actor/lagrangian_loss": logs_dict["lagrangian_loss"].mean(),
"train/rewards_batch": data["rewards"].mean(),
}
if cfg.env.type == "maniskill":
logs.update(
{
"train/return": torch.stack(
[info["log_info"]["return"] for info in infos]
).mean(),
"train/episode_len": torch.stack(
[info["log_info"]["episode_len"] for info in infos]
)
.float()
.mean(),
"train/success": torch.stack(
[info["log_info"]["success"] for info in infos]
)
.float()
.mean(),
}
)
if eval_interval > 0 and global_step % eval_interval == 0:
print(f"Evaluating at global step {global_step}")
if stochastic_eval:
eval_avg_return, eval_avg_length, stoch_eval_info = evaluate(
train_state, stochastic_eval=stochastic_eval
)
eval_avg_return, eval_avg_length, eval_info = evaluate(
train_state
)
eval_info = {
**eval_info,
**{f"stoch/{k}": v for k, v in stoch_eval_info.items()},
}
else:
eval_avg_return, eval_avg_length, eval_info = evaluate(
train_state
)
if cfg.env.type in [
"humanoid_bench",
"isaaclab",
"mtbench",
]:
# NOTE: Hacky way of evaluating performance, but just works
obs, _ = envs.reset()
logs["eval/avg_return"] = eval_avg_return
logs["eval/avg_length"] = eval_avg_length
for key, value in eval_info.items():
if isinstance(value, torch.Tensor):
logs[f"eval/{key}"] = value.mean().item()
elif isinstance(value, np.ndarray):
logs[f"eval/{key}"] = value.mean()
else:
logs[f"eval/{key}"] = value
print(
f"Eval return: {eval_avg_return:.2f}, length: {eval_avg_length:.2f}, env steps: {global_step * cfg.hyperparameters.num_envs * cfg.hyperparameters.num_steps} success rate: {eval_info.get('success', 0.0):.2f}"
)
wandb.log(
{
"speed": speed,
"frame": global_step
* cfg.hyperparameters.num_envs
* cfg.hyperparameters.num_steps,
**logs,
},
step=global_step
* cfg.hyperparameters.num_envs
* cfg.hyperparameters.num_steps,
)
global_step += 1
pbar.update(n=cfg.hyperparameters.num_envs * cfg.hyperparameters.num_steps)
if __name__ == "__main__":
main()
+777
View File
@@ -0,0 +1,777 @@
import os
import torch
import torch.nn as nn
from tensordict import TensorDict
class SimpleReplayBuffer(nn.Module):
def __init__(
self,
n_env: int,
buffer_size: int,
n_obs: int,
n_act: int,
n_critic_obs: int,
asymmetric_obs: bool = False,
playground_mode: bool = False,
n_steps: int = 1,
gamma: float = 0.99,
device=None,
):
"""
A simple replay buffer that stores transitions in a circular buffer.
Supports n-step returns and asymmetric observations.
When playground_mode=True, critic_observations are treated as a concatenation of
regular observations and privileged observations, and only the privileged part is stored
to save memory.
TODO (Younggyo): Refactor to split this into SimpleReplayBuffer and NStepReplayBuffer
"""
super().__init__()
self.n_env = n_env
self.buffer_size = buffer_size
self.n_obs = n_obs
self.n_act = n_act
self.n_critic_obs = n_critic_obs
self.asymmetric_obs = asymmetric_obs
self.playground_mode = playground_mode and asymmetric_obs
self.gamma = gamma
self.n_steps = n_steps
self.device = device
self.observations = torch.zeros(
(n_env, buffer_size, n_obs), device=device, dtype=torch.float
)
self.actions = torch.zeros(
(n_env, buffer_size, n_act), device=device, dtype=torch.float
)
self.rewards = torch.zeros(
(n_env, buffer_size), device=device, dtype=torch.float
)
self.dones = torch.zeros((n_env, buffer_size), device=device, dtype=torch.long)
self.truncations = torch.zeros(
(n_env, buffer_size), device=device, dtype=torch.long
)
self.next_observations = torch.zeros(
(n_env, buffer_size, n_obs), device=device, dtype=torch.float
)
if asymmetric_obs:
if self.playground_mode:
# Only store the privileged part of observations (n_critic_obs - n_obs)
self.privileged_obs_size = n_critic_obs - n_obs
self.privileged_observations = torch.zeros(
(n_env, buffer_size, self.privileged_obs_size),
device=device,
dtype=torch.float,
)
self.next_privileged_observations = torch.zeros(
(n_env, buffer_size, self.privileged_obs_size),
device=device,
dtype=torch.float,
)
else:
# Store full critic observations
self.critic_observations = torch.zeros(
(n_env, buffer_size, n_critic_obs), device=device, dtype=torch.float
)
self.next_critic_observations = torch.zeros(
(n_env, buffer_size, n_critic_obs), device=device, dtype=torch.float
)
self.ptr = 0
def extend(
self,
tensor_dict: TensorDict,
):
observations = tensor_dict["observations"]
actions = tensor_dict["actions"]
rewards = tensor_dict["next"]["rewards"]
dones = tensor_dict["next"]["dones"]
truncations = tensor_dict["next"]["truncations"]
next_observations = tensor_dict["next"]["observations"]
ptr = self.ptr % self.buffer_size
self.observations[:, ptr] = observations
self.actions[:, ptr] = actions
self.rewards[:, ptr] = rewards
self.dones[:, ptr] = dones
self.truncations[:, ptr] = truncations
self.next_observations[:, ptr] = next_observations
if self.asymmetric_obs:
critic_observations = tensor_dict["critic_observations"]
next_critic_observations = tensor_dict["next"]["critic_observations"]
if self.playground_mode:
# Extract and store only the privileged part
privileged_observations = critic_observations[:, self.n_obs :]
next_privileged_observations = next_critic_observations[:, self.n_obs :]
self.privileged_observations[:, ptr] = privileged_observations
self.next_privileged_observations[:, ptr] = next_privileged_observations
else:
# Store full critic observations
self.critic_observations[:, ptr] = critic_observations
self.next_critic_observations[:, ptr] = next_critic_observations
self.ptr += 1
def sample(self, batch_size: int):
# we will sample n_env * batch_size transitions
if self.n_steps == 1:
indices = torch.randint(
0,
min(self.buffer_size, self.ptr),
(self.n_env, batch_size),
device=self.device,
)
obs_indices = indices.unsqueeze(-1).expand(-1, -1, self.n_obs)
act_indices = indices.unsqueeze(-1).expand(-1, -1, self.n_act)
observations = torch.gather(self.observations, 1, obs_indices).reshape(
self.n_env * batch_size, self.n_obs
)
next_observations = torch.gather(
self.next_observations, 1, obs_indices
).reshape(self.n_env * batch_size, self.n_obs)
actions = torch.gather(self.actions, 1, act_indices).reshape(
self.n_env * batch_size, self.n_act
)
rewards = torch.gather(self.rewards, 1, indices).reshape(
self.n_env * batch_size
)
dones = torch.gather(self.dones, 1, indices).reshape(
self.n_env * batch_size
)
truncations = torch.gather(self.truncations, 1, indices).reshape(
self.n_env * batch_size
)
effective_n_steps = torch.ones_like(dones)
if self.asymmetric_obs:
if self.playground_mode:
# Gather privileged observations
priv_obs_indices = indices.unsqueeze(-1).expand(
-1, -1, self.privileged_obs_size
)
privileged_observations = torch.gather(
self.privileged_observations, 1, priv_obs_indices
).reshape(self.n_env * batch_size, self.privileged_obs_size)
next_privileged_observations = torch.gather(
self.next_privileged_observations, 1, priv_obs_indices
).reshape(self.n_env * batch_size, self.privileged_obs_size)
# Concatenate with regular observations to form full critic observations
critic_observations = torch.cat(
[observations, privileged_observations], dim=1
)
next_critic_observations = torch.cat(
[next_observations, next_privileged_observations], dim=1
)
else:
# Gather full critic observations
critic_obs_indices = indices.unsqueeze(-1).expand(
-1, -1, self.n_critic_obs
)
critic_observations = torch.gather(
self.critic_observations, 1, critic_obs_indices
).reshape(self.n_env * batch_size, self.n_critic_obs)
next_critic_observations = torch.gather(
self.next_critic_observations, 1, critic_obs_indices
).reshape(self.n_env * batch_size, self.n_critic_obs)
else:
# Sample base indices
if self.ptr >= self.buffer_size:
# When the buffer is full, there is no protection against sampling across different episodes
# We avoid this by temporarily setting self.pos - 1 to truncated = True if not done
# https://github.com/DLR-RM/stable-baselines3/blob/b91050ca94f8bce7a0285c91f85da518d5a26223/stable_baselines3/common/buffers.py#L857-L860
# TODO (Younggyo): Change the reference when this SB3 branch is merged
current_pos = self.ptr % self.buffer_size
curr_truncations = self.truncations[:, current_pos - 1].clone()
self.truncations[:, current_pos - 1] = torch.logical_not(
self.dones[:, current_pos - 1]
)
indices = torch.randint(
0,
self.buffer_size,
(self.n_env, batch_size),
device=self.device,
)
else:
# Buffer not full - ensure n-step sequence doesn't exceed valid data
max_start_idx = max(1, self.ptr - self.n_steps + 1)
indices = torch.randint(
0,
max_start_idx,
(self.n_env, batch_size),
device=self.device,
)
obs_indices = indices.unsqueeze(-1).expand(-1, -1, self.n_obs)
act_indices = indices.unsqueeze(-1).expand(-1, -1, self.n_act)
# Get base transitions
observations = torch.gather(self.observations, 1, obs_indices).reshape(
self.n_env * batch_size, self.n_obs
)
actions = torch.gather(self.actions, 1, act_indices).reshape(
self.n_env * batch_size, self.n_act
)
if self.asymmetric_obs:
if self.playground_mode:
# Gather privileged observations
priv_obs_indices = indices.unsqueeze(-1).expand(
-1, -1, self.privileged_obs_size
)
privileged_observations = torch.gather(
self.privileged_observations, 1, priv_obs_indices
).reshape(self.n_env * batch_size, self.privileged_obs_size)
# Concatenate with regular observations to form full critic observations
critic_observations = torch.cat(
[observations, privileged_observations], dim=1
)
else:
# Gather full critic observations
critic_obs_indices = indices.unsqueeze(-1).expand(
-1, -1, self.n_critic_obs
)
critic_observations = torch.gather(
self.critic_observations, 1, critic_obs_indices
).reshape(self.n_env * batch_size, self.n_critic_obs)
# Create sequential indices for each sample
# This creates a [n_env, batch_size, n_step] tensor of indices
seq_offsets = torch.arange(self.n_steps, device=self.device).view(1, 1, -1)
all_indices = (
indices.unsqueeze(-1) + seq_offsets
) % self.buffer_size # [n_env, batch_size, n_step]
# Gather all rewards and terminal flags
# Using advanced indexing - result shapes: [n_env, batch_size, n_step]
all_rewards = torch.gather(
self.rewards.unsqueeze(-1).expand(-1, -1, self.n_steps), 1, all_indices
)
all_dones = torch.gather(
self.dones.unsqueeze(-1).expand(-1, -1, self.n_steps), 1, all_indices
)
all_truncations = torch.gather(
self.truncations.unsqueeze(-1).expand(-1, -1, self.n_steps),
1,
all_indices,
)
# Create masks for rewards *after* first done
# This creates a cumulative product that zeroes out rewards after the first done
all_dones_shifted = torch.cat(
[torch.zeros_like(all_dones[:, :, :1]), all_dones[:, :, :-1]], dim=2
) # First reward should not be masked
done_masks = torch.cumprod(
1.0 - all_dones_shifted, dim=2
) # [n_env, batch_size, n_step]
effective_n_steps = done_masks.sum(2)
# Create discount factors
discounts = torch.pow(
self.gamma, torch.arange(self.n_steps, device=self.device)
) # [n_steps]
# Apply masks and discounts to rewards
masked_rewards = all_rewards * done_masks # [n_env, batch_size, n_step]
discounted_rewards = masked_rewards * discounts.view(
1, 1, -1
) # [n_env, batch_size, n_step]
# Sum rewards along the n_step dimension
n_step_rewards = discounted_rewards.sum(dim=2) # [n_env, batch_size]
# Find index of first done or truncation or last step for each sequence
first_done = torch.argmax(
(all_dones > 0).float(), dim=2
) # [n_env, batch_size]
first_trunc = torch.argmax(
(all_truncations > 0).float(), dim=2
) # [n_env, batch_size]
# Handle case where there are no dones or truncations
no_dones = all_dones.sum(dim=2) == 0
no_truncs = all_truncations.sum(dim=2) == 0
# When no dones or truncs, use the last index
first_done = torch.where(no_dones, self.n_steps - 1, first_done)
first_trunc = torch.where(no_truncs, self.n_steps - 1, first_trunc)
# Take the minimum (first) of done or truncation
final_indices = torch.minimum(
first_done, first_trunc
) # [n_env, batch_size]
# Create indices to gather the final next observations
final_next_obs_indices = torch.gather(
all_indices, 2, final_indices.unsqueeze(-1)
).squeeze(-1) # [n_env, batch_size]
# Gather final values
final_next_observations = self.next_observations.gather(
1, final_next_obs_indices.unsqueeze(-1).expand(-1, -1, self.n_obs)
)
final_dones = self.dones.gather(1, final_next_obs_indices)
final_truncations = self.truncations.gather(1, final_next_obs_indices)
if self.asymmetric_obs:
if self.playground_mode:
# Gather final privileged observations
final_next_privileged_observations = (
self.next_privileged_observations.gather(
1,
final_next_obs_indices.unsqueeze(-1).expand(
-1, -1, self.privileged_obs_size
),
)
)
# Reshape for output
next_privileged_observations = (
final_next_privileged_observations.reshape(
self.n_env * batch_size, self.privileged_obs_size
)
)
# Concatenate with next observations to form full next critic observations
next_observations_reshaped = final_next_observations.reshape(
self.n_env * batch_size, self.n_obs
)
next_critic_observations = torch.cat(
[next_observations_reshaped, next_privileged_observations],
dim=1,
)
else:
# Gather final next critic observations directly
final_next_critic_observations = (
self.next_critic_observations.gather(
1,
final_next_obs_indices.unsqueeze(-1).expand(
-1, -1, self.n_critic_obs
),
)
)
next_critic_observations = final_next_critic_observations.reshape(
self.n_env * batch_size, self.n_critic_obs
)
# Reshape everything to batch dimension
rewards = n_step_rewards.reshape(self.n_env * batch_size)
dones = final_dones.reshape(self.n_env * batch_size)
truncations = final_truncations.reshape(self.n_env * batch_size)
effective_n_steps = effective_n_steps.reshape(self.n_env * batch_size)
next_observations = final_next_observations.reshape(
self.n_env * batch_size, self.n_obs
)
out = TensorDict(
{
"observations": observations,
"actions": actions,
"next": {
"rewards": rewards,
"dones": dones,
"truncations": truncations,
"observations": next_observations,
"effective_n_steps": effective_n_steps,
},
},
batch_size=self.n_env * batch_size,
)
if self.asymmetric_obs:
out["critic_observations"] = critic_observations
out["next"]["critic_observations"] = next_critic_observations
if self.n_steps > 1 and self.ptr >= self.buffer_size:
# Roll back the truncation flags introduced for safe sampling
self.truncations[:, current_pos - 1] = curr_truncations
return out
class EmpiricalNormalization(nn.Module):
"""Normalize mean and variance of values based on empirical values."""
def __init__(self, shape, device, eps=1e-2, until=None):
"""Initialize EmpiricalNormalization module.
Args:
shape (int or tuple of int): Shape of input values except batch axis.
eps (float): Small value for stability.
until (int or None): If this arg is specified, the link learns input values until the sum of batch sizes
exceeds it.
"""
super().__init__()
self.eps = eps
self.until = until
self.device = device
self.register_buffer("_mean", torch.zeros(shape).unsqueeze(0).to(device))
self.register_buffer("_var", torch.ones(shape).unsqueeze(0).to(device))
self.register_buffer("_std", torch.ones(shape).unsqueeze(0).to(device))
self.register_buffer("count", torch.tensor(0, dtype=torch.long).to(device))
@property
def mean(self):
return self._mean.squeeze(0).clone()
@property
def std(self):
return self._std.squeeze(0).clone()
def forward(self, x: torch.Tensor, center: bool = True) -> torch.Tensor:
if x.shape[-1:] != self._mean.shape[-1:]:
raise ValueError(
f"Expected input of shape (*,{self._mean.shape[-1:]}), got {x.shape}"
)
if self.training:
self.update(x)
if center:
return (x - self._mean) / (self._std + self.eps)
else:
return x / (self._std + self.eps)
@torch.jit.unused
def update(self, x):
x = x.flatten(end_dim=-2)
if self.until is not None and self.count >= self.until:
return
batch_size = x.shape[0]
batch_mean = torch.mean(x, dim=0, keepdim=True)
# Update count
new_count = self.count + batch_size
# Update mean
delta = batch_mean - self._mean
self._mean += (batch_size / new_count) * delta
# Update variance using Chan's parallel algorithm
# https://en.wikipedia.org/wiki/Algorithms_for_calculating_variance#Parallel_algorithm
if self.count > 0: # Ensure we're not dividing by zero
batch_var = torch.mean((x - batch_mean) ** 2, dim=0, keepdim=True)
delta2 = batch_mean - self._mean
m_a = self._var * self.count
m_b = batch_var * batch_size
M2 = m_a + m_b + (delta2**2) * (self.count * batch_size / new_count)
self._var = M2 / new_count
else:
# For first batch, just use batch variance
self._var = torch.mean((x - self._mean) ** 2, dim=0, keepdim=True)
self._std = torch.sqrt(self._var)
self.count = new_count
@torch.jit.unused
def inverse(self, y):
return y * (self._std + self.eps) + self._mean
class RewardNormalizer(nn.Module):
def __init__(
self,
gamma: float,
device: torch.device,
g_max: float = 10.0,
epsilon: float = 1e-8,
):
super().__init__()
self.register_buffer(
"G", torch.zeros(1, device=device)
) # running estimate of the discounted return
self.register_buffer("G_r_max", torch.zeros(1, device=device)) # running-max
self.G_rms = EmpiricalNormalization(shape=1, device=device)
self.gamma = gamma
self.g_max = g_max
self.epsilon = epsilon
def _scale_reward(self, rewards: torch.Tensor) -> torch.Tensor:
var_denominator = self.G_rms.std[0] + self.epsilon
min_required_denominator = self.G_r_max / self.g_max
denominator = torch.maximum(var_denominator, min_required_denominator)
return rewards / denominator
def update_stats(
self,
rewards: torch.Tensor,
dones: torch.Tensor,
):
self.G = self.gamma * (1 - dones) * self.G + rewards
self.G_rms.update(self.G.view(-1, 1))
self.G_r_max = max(self.G_r_max, max(abs(self.G)))
def forward(self, rewards: torch.Tensor) -> torch.Tensor:
return self._scale_reward(rewards)
class PerTaskEmpiricalNormalization(nn.Module):
"""Normalize mean and variance of values based on empirical values for each task."""
def __init__(
self,
num_tasks: int,
shape: tuple,
device: torch.device,
eps: float = 1e-2,
until: int = None,
):
"""
Initialize PerTaskEmpiricalNormalization module.
Args:
num_tasks (int): The total number of tasks.
shape (int or tuple of int): Shape of input values except batch axis.
eps (float): Small value for stability.
until (int or None): If specified, learns until the sum of batch sizes
for a specific task exceeds this value.
"""
super().__init__()
if not isinstance(shape, tuple):
shape = (shape,)
self.num_tasks = num_tasks
self.shape = shape
self.eps = eps
self.until = until
self.device = device
# Buffers now have a leading dimension for tasks
self.register_buffer("_mean", torch.zeros(num_tasks, *shape).to(device))
self.register_buffer("_var", torch.ones(num_tasks, *shape).to(device))
self.register_buffer("_std", torch.ones(num_tasks, *shape).to(device))
self.register_buffer(
"count", torch.zeros(num_tasks, dtype=torch.long).to(device)
)
def forward(
self, x: torch.Tensor, task_ids: torch.Tensor, center: bool = True
) -> torch.Tensor:
"""
Normalize the input tensor `x` using statistics for the given `task_ids`.
Args:
x (torch.Tensor): Input tensor of shape [num_envs, *shape].
task_ids (torch.Tensor): Tensor of task indices, shape [num_envs].
center (bool): If True, center the data by subtracting the mean.
"""
if x.shape[1:] != self.shape:
raise ValueError(f"Expected input shape (*, {self.shape}), got {x.shape}")
if x.shape[0] != task_ids.shape[0]:
raise ValueError("Batch size of x and task_ids must match.")
# Gather the stats for the tasks in the current batch
# Reshape task_ids for broadcasting: [num_envs] -> [num_envs, 1, ...]
view_shape = (task_ids.shape[0],) + (1,) * len(self.shape)
task_ids_expanded = task_ids.view(view_shape).expand_as(x)
mean = self._mean.gather(0, task_ids_expanded)
std = self._std.gather(0, task_ids_expanded)
if self.training:
self.update(x, task_ids)
if center:
return (x - mean) / (std + self.eps)
else:
return x / (std + self.eps)
@torch.jit.unused
def update(self, x: torch.Tensor, task_ids: torch.Tensor):
"""Update running statistics for the tasks present in the batch."""
unique_tasks = torch.unique(task_ids)
for task_id in unique_tasks:
if self.until is not None and self.count[task_id] >= self.until:
continue
# Create a mask to select data for the current task
mask = task_ids == task_id
x_task = x[mask]
batch_size = x_task.shape[0]
if batch_size == 0:
continue
# Update count for this task
old_count = self.count[task_id].clone()
new_count = old_count + batch_size
# Update mean
task_mean = self._mean[task_id]
batch_mean = torch.mean(x_task, dim=0)
delta = batch_mean - task_mean
self._mean[task_id] = task_mean + (batch_size / new_count) * delta
# Update variance using Chan's parallel algorithm
if old_count > 0:
batch_var = torch.var(x_task, dim=0, unbiased=False)
m_a = self._var[task_id] * old_count
m_b = batch_var * batch_size
M2 = m_a + m_b + (delta**2) * (old_count * batch_size / new_count)
self._var[task_id] = M2 / new_count
else:
# For the first batch of this task
self._var[task_id] = torch.var(x_task, dim=0, unbiased=False)
self._std[task_id] = torch.sqrt(self._var[task_id])
self.count[task_id] = new_count
class PerTaskRewardNormalizer(nn.Module):
def __init__(
self,
num_tasks: int,
gamma: float,
device: torch.device,
g_max: float = 10.0,
epsilon: float = 1e-8,
):
"""
Per-task reward normalizer, motivation comes from BRC (https://arxiv.org/abs/2505.23150v1)
"""
super().__init__()
self.num_tasks = num_tasks
self.gamma = gamma
self.g_max = g_max
self.epsilon = epsilon
self.device = device
# Per-task running estimate of the discounted return
self.register_buffer("G", torch.zeros(num_tasks, device=device))
# Per-task running-max of the discounted return
self.register_buffer("G_r_max", torch.zeros(num_tasks, device=device))
# Use the new per-task normalizer for the statistics of G
self.G_rms = PerTaskEmpiricalNormalization(
num_tasks=num_tasks, shape=(1,), device=device
)
def _scale_reward(
self, rewards: torch.Tensor, task_ids: torch.Tensor
) -> torch.Tensor:
"""
Scales rewards using per-task statistics.
Args:
rewards (torch.Tensor): Reward tensor, shape [num_envs].
task_ids (torch.Tensor): Task indices, shape [num_envs].
"""
# Gather stats for the tasks in the batch
std_for_batch = self.G_rms._std.gather(0, task_ids.unsqueeze(-1)).squeeze(-1)
g_r_max_for_batch = self.G_r_max.gather(0, task_ids)
var_denominator = std_for_batch + self.epsilon
min_required_denominator = g_r_max_for_batch / self.g_max
denominator = torch.maximum(var_denominator, min_required_denominator)
# Add a small epsilon to the final denominator to prevent division by zero
# in case g_r_max is also zero.
return rewards / (denominator + self.epsilon)
def update_stats(
self, rewards: torch.Tensor, dones: torch.Tensor, task_ids: torch.Tensor
):
"""
Updates the running discounted return and its statistics for each task.
Args:
rewards (torch.Tensor): Reward tensor, shape [num_envs].
dones (torch.Tensor): Done tensor, shape [num_envs].
task_ids (torch.Tensor): Task indices, shape [num_envs].
"""
if not (rewards.shape == dones.shape == task_ids.shape):
raise ValueError("rewards, dones, and task_ids must have the same shape.")
# === Update G (running discounted return) ===
# Gather the previous G values for the tasks in the batch
prev_G = self.G.gather(0, task_ids)
# Update G for each environment based on its own reward and done signal
new_G = self.gamma * (1 - dones.float()) * prev_G + rewards
# Scatter the updated G values back to the main buffer
self.G.scatter_(0, task_ids, new_G)
# === Update G_rms (statistics of G) ===
# The update function handles the per-task logic internally
self.G_rms.update(new_G.unsqueeze(-1), task_ids)
# === Update G_r_max (running max of |G|) ===
prev_G_r_max = self.G_r_max.gather(0, task_ids)
# Update the max for each environment
updated_G_r_max = torch.maximum(prev_G_r_max, torch.abs(new_G))
# Scatter the new maxes back to the main buffer
self.G_r_max.scatter_(0, task_ids, updated_G_r_max)
def forward(self, rewards: torch.Tensor, task_ids: torch.Tensor) -> torch.Tensor:
"""
Normalizes rewards. During training, it also updates the running statistics.
Args:
rewards (torch.Tensor): Reward tensor, shape [num_envs].
task_ids (torch.Tensor): Task indices, shape [num_envs].
"""
return self._scale_reward(rewards, task_ids)
def cpu_state(sd):
# detach & move to host without locking the compute stream
return {k: v.detach().to("cpu", non_blocking=True) for k, v in sd.items()}
def save_params(
global_step,
actor,
qnet,
qnet_target,
obs_normalizer,
critic_obs_normalizer,
args,
save_path,
):
"""Save model parameters and training configuration to disk."""
os.makedirs(os.path.dirname(save_path), exist_ok=True)
save_dict = {
"actor_state_dict": cpu_state(actor.state_dict()),
"qnet_state_dict": cpu_state(qnet.state_dict()),
"qnet_target_state_dict": cpu_state(qnet_target.state_dict()),
"obs_normalizer_state": (
cpu_state(obs_normalizer.state_dict())
if hasattr(obs_normalizer, "state_dict")
else None
),
"critic_obs_normalizer_state": (
cpu_state(critic_obs_normalizer.state_dict())
if hasattr(critic_obs_normalizer, "state_dict")
else None
),
"args": vars(args), # Save all arguments
"global_step": global_step,
}
torch.save(save_dict, save_path, _use_new_zipfile_serialization=True)
print(f"Saved parameters and configuration to {save_path}")
def hl_gauss(inp, vmin, vmax, num_atoms):
x = torch.clip(inp, vmin, max=vmax)
bin_width = (vmax - vmin) / (num_atoms - 1)
sigma_to_final_sigma_ratio = 0.75
support = torch.linspace(
vmin - bin_width / 2,
vmax + bin_width / 2,
num_atoms + 1,
device=inp.device,
)
sigma = bin_width * sigma_to_final_sigma_ratio
cdf_evals = torch.erf(
(support.unsqueeze(0) - x).squeeze()
/ (torch.sqrt(torch.tensor(2.0)) * sigma + 1e-6)
)
z = cdf_evals[..., -1] - cdf_evals[..., 0]
target_probs = cdf_evals[..., 1:] - cdf_evals[..., :-1]
target_probs = (target_probs / (z.unsqueeze(-1) + 1e-6)).reshape(
*inp.shape[:-1], num_atoms
)
return target_probs