Fixes build errors due to name conflicts
This commit is contained in:
@@ -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'."
|
||||
)
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user