release
This commit is contained in:
@@ -0,0 +1,170 @@
|
||||
"""
|
||||
Parent fine-tuning agent class.
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
import numpy as np
|
||||
from omegaconf import OmegaConf
|
||||
import torch
|
||||
import hydra
|
||||
import logging
|
||||
import wandb
|
||||
import random
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from env.gym_utils import make_async
|
||||
|
||||
|
||||
class TrainAgent:
|
||||
|
||||
def __init__(self, cfg):
|
||||
super().__init__()
|
||||
self.cfg = cfg
|
||||
self.device = cfg.device
|
||||
self.seed = cfg.get("seed", 42)
|
||||
random.seed(self.seed)
|
||||
np.random.seed(self.seed)
|
||||
torch.manual_seed(self.seed)
|
||||
|
||||
# Wandb
|
||||
self.use_wandb = cfg.wandb is not None
|
||||
if cfg.wandb is not None:
|
||||
wandb.init(
|
||||
entity=cfg.wandb.entity,
|
||||
project=cfg.wandb.project,
|
||||
name=cfg.wandb.run,
|
||||
config=OmegaConf.to_container(cfg, resolve=True),
|
||||
)
|
||||
|
||||
# Make vectorized env
|
||||
self.env_name = cfg.env.name
|
||||
env_type = cfg.env.get("env_type", None)
|
||||
self.venv = make_async(
|
||||
cfg.env.name,
|
||||
env_type=env_type,
|
||||
num_envs=cfg.env.n_envs,
|
||||
asynchronous=True,
|
||||
max_episode_steps=cfg.env.max_episode_steps,
|
||||
wrappers=cfg.env.get("wrappers", None),
|
||||
robomimic_env_cfg_path=cfg.get("robomimic_env_cfg_path", None),
|
||||
shape_meta=cfg.get("shape_meta", None),
|
||||
use_image_obs=cfg.env.get("use_image_obs", False),
|
||||
render=cfg.env.get("render", False),
|
||||
render_offscreen=cfg.env.get("save_video", False),
|
||||
obs_dim=cfg.obs_dim,
|
||||
action_dim=cfg.action_dim,
|
||||
**cfg.env.specific if "specific" in cfg.env else {},
|
||||
)
|
||||
if not env_type == "furniture":
|
||||
self.venv.seed(
|
||||
[self.seed + i for i in range(cfg.env.n_envs)]
|
||||
) # otherwise parallel envs might have the same initial states!
|
||||
# isaacgym environments do not need seeding
|
||||
self.n_envs = cfg.env.n_envs
|
||||
self.n_cond_step = cfg.cond_steps
|
||||
self.obs_dim = cfg.obs_dim
|
||||
self.action_dim = cfg.action_dim
|
||||
self.act_steps = cfg.act_steps
|
||||
self.horizon_steps = cfg.horizon_steps
|
||||
self.max_episode_steps = cfg.env.max_episode_steps
|
||||
self.reset_at_iteration = cfg.env.get("reset_at_iteration", True)
|
||||
self.save_full_observations = cfg.env.get("save_full_observations", False)
|
||||
self.furniture_sparse_reward = (
|
||||
cfg.env.specific.get("sparse_reward", False)
|
||||
if "specific" in cfg.env
|
||||
else False
|
||||
) # furniture specific, for best reward calculation
|
||||
|
||||
# Batch size for gradient update
|
||||
self.batch_size: int = cfg.train.batch_size
|
||||
|
||||
# Build model and load checkpoint
|
||||
self.model = hydra.utils.instantiate(cfg.model)
|
||||
|
||||
# Training params
|
||||
self.itr = 0
|
||||
self.n_train_itr = cfg.train.n_train_itr
|
||||
self.val_freq = cfg.train.val_freq
|
||||
self.force_train = cfg.train.get("force_train", False)
|
||||
self.n_steps = cfg.train.n_steps
|
||||
self.best_reward_threshold_for_success = (
|
||||
len(self.venv.pairs_to_assemble)
|
||||
if env_type == "furniture"
|
||||
else cfg.env.best_reward_threshold_for_success
|
||||
)
|
||||
self.max_grad_norm = cfg.train.get("max_grad_norm", None)
|
||||
|
||||
# Logging, rendering, checkpoints
|
||||
self.logdir = cfg.logdir
|
||||
self.render_dir = os.path.join(self.logdir, "render")
|
||||
self.checkpoint_dir = os.path.join(self.logdir, "checkpoint")
|
||||
self.result_path = os.path.join(self.logdir, "result.pkl")
|
||||
os.makedirs(self.render_dir, exist_ok=True)
|
||||
os.makedirs(self.checkpoint_dir, exist_ok=True)
|
||||
self.save_trajs = cfg.train.get("save_trajs", False)
|
||||
self.log_freq = cfg.train.get("log_freq", 1)
|
||||
self.save_model_freq = cfg.train.save_model_freq
|
||||
self.render_freq = cfg.train.render.freq
|
||||
self.n_render = cfg.train.render.num
|
||||
self.render_video = cfg.env.get("save_video", False)
|
||||
assert self.n_render <= self.n_envs, "n_render must be <= n_envs"
|
||||
assert not (
|
||||
self.n_render <= 0 and self.render_video
|
||||
), "Need to set n_render > 0 if saving video"
|
||||
self.traj_plotter = (
|
||||
hydra.utils.instantiate(cfg.train.plotter)
|
||||
if "plotter" in cfg.train
|
||||
else None
|
||||
)
|
||||
|
||||
def run(self):
|
||||
pass
|
||||
|
||||
def save_model(self):
|
||||
"""
|
||||
saves model to disk; no ema
|
||||
"""
|
||||
data = {
|
||||
"itr": self.itr,
|
||||
"model": self.model.state_dict(),
|
||||
}
|
||||
savepath = os.path.join(self.checkpoint_dir, f"state_{self.itr}.pt")
|
||||
torch.save(data, savepath)
|
||||
log.info(f"Saved model to {savepath}")
|
||||
|
||||
def load(self, itr):
|
||||
"""
|
||||
loads model from disk
|
||||
"""
|
||||
loadpath = os.path.join(self.checkpoint_dir, f"state_{itr}.pt")
|
||||
data = torch.load(loadpath, weights_only=True)
|
||||
|
||||
self.itr = data["itr"]
|
||||
self.model.load_state_dict(data["model"])
|
||||
|
||||
def reset_env_all(self, verbose=False, options_venv=None, **kwargs):
|
||||
if options_venv is None:
|
||||
options_venv = [
|
||||
{k: v for k, v in kwargs.items()} for _ in range(self.n_envs)
|
||||
]
|
||||
obs_venv = self.venv.reset_arg(options_list=options_venv)
|
||||
# convert to OrderedDict if obs_venv is a list of dict
|
||||
if isinstance(obs_venv, list):
|
||||
obs_venv = {
|
||||
key: np.stack([obs_venv[i][key] for i in range(self.n_envs)])
|
||||
for key in obs_venv[0].keys()
|
||||
}
|
||||
if verbose:
|
||||
for index in range(self.n_envs):
|
||||
logging.info(
|
||||
f"<-- Reset environment {index} with options {options_venv[index]}"
|
||||
)
|
||||
return obs_venv
|
||||
|
||||
def reset_env(self, env_ind, verbose=False):
|
||||
task = {}
|
||||
obs = self.venv.reset_one_arg(env_ind=env_ind, options=task)
|
||||
if verbose:
|
||||
logging.info(f"<-- Reset environment {env_ind} with task {task}")
|
||||
return obs
|
||||
@@ -0,0 +1,389 @@
|
||||
"""
|
||||
Advantage-weighted regression (AWR) for diffusion policy.
|
||||
|
||||
Advantage = discounted-reward-to-go - V(s)
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
import pickle
|
||||
import einops
|
||||
import numpy as np
|
||||
import torch
|
||||
import logging
|
||||
import wandb
|
||||
from copy import deepcopy
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from util.timer import Timer
|
||||
from collections import deque
|
||||
from agent.finetune.train_agent import TrainAgent
|
||||
from util.scheduler import CosineAnnealingWarmupRestarts
|
||||
|
||||
|
||||
def td_values(
|
||||
states,
|
||||
rewards,
|
||||
dones,
|
||||
state_values,
|
||||
gamma=0.99,
|
||||
alpha=0.95,
|
||||
lam=0.95,
|
||||
):
|
||||
"""
|
||||
Gives a list of TD estimates for a given list of samples from an RL environment.
|
||||
The TD(λ) estimator is used for this computation.
|
||||
|
||||
:param replay_buffers: The replay buffers filled by exploring the RL environment.
|
||||
Includes: states, rewards, "final state?"s.
|
||||
:param state_values: The currently estimated state values.
|
||||
:return: The TD estimates.
|
||||
"""
|
||||
sample_count = len(states)
|
||||
tds = np.zeros_like(state_values, dtype=np.float32)
|
||||
dones[-1] = 1
|
||||
next_value = 1 - dones[-1]
|
||||
|
||||
val = 0.0
|
||||
for i in range(sample_count - 1, -1, -1):
|
||||
# next_value = 0.0 if dones[i] else state_values[i + 1]
|
||||
|
||||
# get next_value for vectorized
|
||||
if i < sample_count - 1:
|
||||
next_value = state_values[i + 1]
|
||||
next_value = next_value * (1 - dones[i])
|
||||
|
||||
state_value = state_values[i]
|
||||
error = rewards[i] + gamma * next_value - state_value
|
||||
val = alpha * error + gamma * lam * (1 - dones[i]) * val
|
||||
|
||||
tds[i] = val + state_value
|
||||
return tds
|
||||
|
||||
|
||||
class TrainAWRDiffusionAgent(TrainAgent):
|
||||
|
||||
def __init__(self, cfg):
|
||||
super().__init__(cfg)
|
||||
self.logprob_batch_size = cfg.train.get("logprob_batch_size", 10000)
|
||||
|
||||
# note the discount factor gamma here is applied to reward every act_steps, instead of every env step
|
||||
self.gamma = cfg.train.gamma
|
||||
|
||||
# Wwarm up period for critic before actor updates
|
||||
self.n_critic_warmup_itr = cfg.train.n_critic_warmup_itr
|
||||
|
||||
# Optimizer
|
||||
self.actor_optimizer = torch.optim.AdamW(
|
||||
self.model.actor.parameters(),
|
||||
lr=cfg.train.actor_lr,
|
||||
weight_decay=cfg.train.actor_weight_decay,
|
||||
)
|
||||
self.actor_lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.actor_optimizer,
|
||||
first_cycle_steps=cfg.train.actor_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.actor_lr,
|
||||
min_lr=cfg.train.actor_lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.actor_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
self.critic_optimizer = torch.optim.AdamW(
|
||||
self.model.critic.parameters(),
|
||||
lr=cfg.train.critic_lr,
|
||||
weight_decay=cfg.train.critic_weight_decay,
|
||||
)
|
||||
self.critic_lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.critic_optimizer,
|
||||
first_cycle_steps=cfg.train.critic_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.critic_lr,
|
||||
min_lr=cfg.train.critic_lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.critic_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
|
||||
# Buffer size
|
||||
self.buffer_size = cfg.train.buffer_size
|
||||
|
||||
# Reward exponential
|
||||
self.beta = cfg.train.beta
|
||||
|
||||
# Max weight for AWR
|
||||
self.max_adv_weight = cfg.train.max_adv_weight
|
||||
|
||||
# Scaling reward
|
||||
self.scale_reward_factor = cfg.train.scale_reward_factor
|
||||
|
||||
# Updates
|
||||
self.replay_ratio = cfg.train.replay_ratio
|
||||
self.critic_update_ratio = cfg.train.critic_update_ratio
|
||||
|
||||
def run(self):
|
||||
|
||||
# make a FIFO replay buffer for obs, action, and reward
|
||||
obs_buffer = deque(maxlen=self.buffer_size)
|
||||
action_buffer = deque(maxlen=self.buffer_size)
|
||||
reward_buffer = deque(maxlen=self.buffer_size)
|
||||
done_buffer = deque(maxlen=self.buffer_size)
|
||||
first_buffer = deque(maxlen=self.buffer_size)
|
||||
|
||||
# Start training loop
|
||||
timer = Timer()
|
||||
run_results = []
|
||||
done_venv = np.zeros((1, self.n_envs))
|
||||
while self.itr < self.n_train_itr:
|
||||
|
||||
# Prepare video paths for each envs --- only applies for the first set of episodes if allowing reset within iteration and each iteration has multiple episodes from one env
|
||||
options_venv = [{} for _ in range(self.n_envs)]
|
||||
if self.itr % self.render_freq == 0 and self.render_video:
|
||||
for env_ind in range(self.n_render):
|
||||
options_venv[env_ind]["video_path"] = os.path.join(
|
||||
self.render_dir, f"itr-{self.itr}_trial-{env_ind}.mp4"
|
||||
)
|
||||
|
||||
# Define train or eval - all envs restart
|
||||
eval_mode = self.itr % self.val_freq == 0 and not self.force_train
|
||||
self.model.eval() if eval_mode else self.model.train()
|
||||
firsts_trajs = np.zeros((self.n_steps + 1, self.n_envs))
|
||||
|
||||
# Reset env at the beginning of an iteration
|
||||
if self.reset_at_iteration or eval_mode or last_itr_eval:
|
||||
prev_obs_venv = self.reset_env_all(options_venv=options_venv)
|
||||
firsts_trajs[0] = 1
|
||||
else:
|
||||
firsts_trajs[0] = (
|
||||
done_venv # if done at the end of last iteration, then the envs are just reset
|
||||
)
|
||||
last_itr_eval = eval_mode
|
||||
reward_trajs = np.empty((0, self.n_envs))
|
||||
|
||||
# Collect a set of trajectories from env
|
||||
for step in range(self.n_steps):
|
||||
if step % 10 == 0:
|
||||
print(f"Processed step {step} of {self.n_steps}")
|
||||
|
||||
# Select action
|
||||
with torch.no_grad():
|
||||
samples = (
|
||||
self.model(
|
||||
cond=torch.from_numpy(prev_obs_venv)
|
||||
.float()
|
||||
.to(self.device),
|
||||
deterministic=eval_mode,
|
||||
)
|
||||
.cpu()
|
||||
.numpy()
|
||||
) # n_env x horizon x act
|
||||
action_venv = samples[:, : self.act_steps]
|
||||
|
||||
# Apply multi-step action
|
||||
obs_venv, reward_venv, done_venv, info_venv = self.venv.step(
|
||||
action_venv
|
||||
)
|
||||
reward_trajs = np.vstack((reward_trajs, reward_venv[None]))
|
||||
|
||||
# add to buffer
|
||||
obs_buffer.append(prev_obs_venv)
|
||||
action_buffer.append(action_venv)
|
||||
reward_buffer.append(reward_venv * self.scale_reward_factor)
|
||||
done_buffer.append(done_venv)
|
||||
first_buffer.append(firsts_trajs[step])
|
||||
|
||||
firsts_trajs[step + 1] = done_venv
|
||||
prev_obs_venv = obs_venv
|
||||
|
||||
# Summarize episode reward --- this needs to be handled differently depending on whether the environment is reset after each iteration. Only count episodes that finish within the iteration.
|
||||
episodes_start_end = []
|
||||
for env_ind in range(self.n_envs):
|
||||
env_steps = np.where(firsts_trajs[:, env_ind] == 1)[0]
|
||||
for i in range(len(env_steps) - 1):
|
||||
start = env_steps[i]
|
||||
end = env_steps[i + 1]
|
||||
if end - start > 1:
|
||||
episodes_start_end.append((env_ind, start, end - 1))
|
||||
if len(episodes_start_end) > 0:
|
||||
reward_trajs_split = [
|
||||
reward_trajs[start : end + 1, env_ind]
|
||||
for env_ind, start, end in episodes_start_end
|
||||
]
|
||||
num_episode_finished = len(reward_trajs_split)
|
||||
episode_reward = np.array(
|
||||
[np.sum(reward_traj) for reward_traj in reward_trajs_split]
|
||||
)
|
||||
episode_best_reward = np.array(
|
||||
[
|
||||
np.max(reward_traj) / self.act_steps
|
||||
for reward_traj in reward_trajs_split
|
||||
]
|
||||
)
|
||||
avg_episode_reward = np.mean(episode_reward)
|
||||
avg_best_reward = np.mean(episode_best_reward)
|
||||
success_rate = np.mean(
|
||||
episode_best_reward >= self.best_reward_threshold_for_success
|
||||
)
|
||||
else:
|
||||
episode_reward = np.array([])
|
||||
num_episode_finished = 0
|
||||
avg_episode_reward = 0
|
||||
avg_best_reward = 0
|
||||
success_rate = 0
|
||||
log.info("[WARNING] No episode completed within the iteration!")
|
||||
|
||||
# Update
|
||||
if not eval_mode:
|
||||
|
||||
obs_trajs = np.array(deepcopy(obs_buffer))
|
||||
reward_trajs = np.array(deepcopy(reward_buffer))
|
||||
dones_trajs = np.array(deepcopy(done_buffer))
|
||||
|
||||
obs_t = einops.rearrange(
|
||||
torch.from_numpy(obs_trajs).float().to(self.device),
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
values_t = np.array(self.model.critic(obs_t).detach().cpu().numpy())
|
||||
values_trajs = values_t.reshape(-1, self.n_envs)
|
||||
td_trajs = td_values(obs_trajs, reward_trajs, dones_trajs, values_trajs)
|
||||
|
||||
# flatten
|
||||
obs_trajs = einops.rearrange(
|
||||
obs_trajs,
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
td_trajs = einops.rearrange(
|
||||
td_trajs,
|
||||
"s e -> (s e)",
|
||||
)
|
||||
|
||||
# Update policy and critic
|
||||
num_batch = int(
|
||||
self.n_steps * self.n_envs / self.batch_size * self.replay_ratio
|
||||
)
|
||||
for _ in range(num_batch // self.critic_update_ratio):
|
||||
|
||||
# Sample batch
|
||||
inds = np.random.choice(len(obs_trajs), self.batch_size)
|
||||
obs_b = torch.from_numpy(obs_trajs[inds]).float().to(self.device)
|
||||
td_b = torch.from_numpy(td_trajs[inds]).float().to(self.device)
|
||||
|
||||
# Update critic
|
||||
loss_critic = self.model.loss_critic(obs_b, td_b)
|
||||
self.critic_optimizer.zero_grad()
|
||||
loss_critic.backward()
|
||||
self.critic_optimizer.step()
|
||||
|
||||
obs_trajs = np.array(deepcopy(obs_buffer))
|
||||
samples_trajs = np.array(deepcopy(action_buffer))
|
||||
reward_trajs = np.array(deepcopy(reward_buffer))
|
||||
dones_trajs = np.array(deepcopy(done_buffer))
|
||||
|
||||
obs_t = einops.rearrange(
|
||||
torch.from_numpy(obs_trajs).float().to(self.device),
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
values_t = np.array(self.model.critic(obs_t).detach().cpu().numpy())
|
||||
values_trajs = values_t.reshape(-1, self.n_envs)
|
||||
td_trajs = td_values(obs_trajs, reward_trajs, dones_trajs, values_trajs)
|
||||
advantages_trajs = td_trajs - values_trajs
|
||||
|
||||
# flatten
|
||||
obs_trajs = einops.rearrange(
|
||||
obs_trajs,
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
samples_trajs = einops.rearrange(
|
||||
samples_trajs,
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
advantages_trajs = einops.rearrange(
|
||||
advantages_trajs,
|
||||
"s e -> (s e)",
|
||||
)
|
||||
|
||||
for _ in range(num_batch):
|
||||
|
||||
# Sample batch
|
||||
inds = np.random.choice(len(obs_trajs), self.batch_size)
|
||||
obs_b = torch.from_numpy(obs_trajs[inds]).float().to(self.device)
|
||||
actions_b = (
|
||||
torch.from_numpy(samples_trajs[inds]).float().to(self.device)
|
||||
)
|
||||
advantages_b = (
|
||||
torch.from_numpy(advantages_trajs[inds]).float().to(self.device)
|
||||
)
|
||||
advantages_b = (advantages_b - advantages_b.mean()) / (
|
||||
advantages_b.std() + 1e-6
|
||||
)
|
||||
advantages_b_scaled = torch.exp(self.beta * advantages_b)
|
||||
advantages_b_scaled.clamp_(max=self.max_adv_weight)
|
||||
|
||||
# Update policy with collected trajectories
|
||||
loss = self.model.loss(
|
||||
actions_b,
|
||||
obs_b,
|
||||
advantages_b_scaled.detach(),
|
||||
)
|
||||
self.actor_optimizer.zero_grad()
|
||||
loss.backward()
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
if self.max_grad_norm is not None:
|
||||
torch.nn.utils.clip_grad_norm_(
|
||||
self.model.actor.parameters(), self.max_grad_norm
|
||||
)
|
||||
self.actor_optimizer.step()
|
||||
|
||||
# Update lr
|
||||
self.actor_lr_scheduler.step()
|
||||
self.critic_lr_scheduler.step()
|
||||
|
||||
# Save model
|
||||
if self.itr % self.save_model_freq == 0 or self.itr == self.n_train_itr - 1:
|
||||
self.save_model()
|
||||
|
||||
# Log loss and save metrics
|
||||
run_results.append(
|
||||
{
|
||||
"itr": self.itr,
|
||||
}
|
||||
)
|
||||
if self.itr % self.log_freq == 0:
|
||||
if eval_mode:
|
||||
log.info(
|
||||
f"eval: success rate {success_rate:8.4f} | avg episode reward {avg_episode_reward:8.4f} | avg best reward {avg_best_reward:8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"success rate - eval": success_rate,
|
||||
"avg episode reward - eval": avg_episode_reward,
|
||||
"avg best reward - eval": avg_best_reward,
|
||||
"num episode - eval": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=False,
|
||||
)
|
||||
run_results[-1]["eval_success_rate"] = success_rate
|
||||
run_results[-1]["eval_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["eval_best_reward"] = avg_best_reward
|
||||
else:
|
||||
log.info(
|
||||
f"{self.itr}: loss {loss:8.4f} | reward {avg_episode_reward:8.4f} |t:{timer():8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"loss": loss,
|
||||
"loss - critic": loss_critic,
|
||||
"avg episode reward - train": avg_episode_reward,
|
||||
"num episode - train": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=True,
|
||||
)
|
||||
run_results[-1]["loss"] = loss
|
||||
run_results[-1]["loss_critic"] = loss_critic
|
||||
run_results[-1]["train_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["time"] = timer()
|
||||
with open(self.result_path, "wb") as f:
|
||||
pickle.dump(run_results, f)
|
||||
self.itr += 1
|
||||
@@ -0,0 +1,358 @@
|
||||
"""
|
||||
Model-free online RL with DIffusion POlicy (DIPO)
|
||||
|
||||
Applies action gradient to perturb actions towards maximizer of Q-function.
|
||||
|
||||
a_t <- a_t + \eta * \grad_a Q(s, a)
|
||||
"""
|
||||
|
||||
import os
|
||||
import pickle
|
||||
import numpy as np
|
||||
import torch
|
||||
import logging
|
||||
import wandb
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from util.timer import Timer
|
||||
from collections import deque
|
||||
from agent.finetune.train_agent import TrainAgent
|
||||
from util.scheduler import CosineAnnealingWarmupRestarts
|
||||
|
||||
|
||||
class TrainDIPODiffusionAgent(TrainAgent):
|
||||
|
||||
def __init__(self, cfg):
|
||||
super().__init__(cfg)
|
||||
|
||||
# note the discount factor gamma here is applied to reward every act_steps, instead of every env step
|
||||
self.gamma = cfg.train.gamma
|
||||
|
||||
# Wwarm up period for critic before actor updates
|
||||
self.n_critic_warmup_itr = cfg.train.n_critic_warmup_itr
|
||||
|
||||
# Optimizer
|
||||
self.actor_optimizer = torch.optim.AdamW(
|
||||
self.model.actor.parameters(),
|
||||
lr=cfg.train.actor_lr,
|
||||
weight_decay=cfg.train.actor_weight_decay,
|
||||
)
|
||||
# use cosine scheduler with linear warmup
|
||||
self.actor_lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.actor_optimizer,
|
||||
first_cycle_steps=cfg.train.actor_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.actor_lr,
|
||||
min_lr=cfg.train.actor_lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.actor_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
self.critic_optimizer = torch.optim.AdamW(
|
||||
self.model.critic.parameters(),
|
||||
lr=cfg.train.critic_lr,
|
||||
weight_decay=cfg.train.critic_weight_decay,
|
||||
)
|
||||
self.critic_lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.critic_optimizer,
|
||||
first_cycle_steps=cfg.train.critic_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.critic_lr,
|
||||
min_lr=cfg.train.critic_lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.critic_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
|
||||
# Buffer size
|
||||
self.buffer_size = cfg.train.buffer_size
|
||||
|
||||
# Perturbation scale
|
||||
self.eta = cfg.train.eta
|
||||
|
||||
# Updates
|
||||
self.replay_ratio = cfg.train.replay_ratio
|
||||
|
||||
# Scaling reward
|
||||
self.scale_reward_factor = cfg.train.scale_reward_factor
|
||||
|
||||
# Apply action gradient many steps
|
||||
self.action_gradient_steps = cfg.train.action_gradient_steps
|
||||
|
||||
def run(self):
|
||||
|
||||
# make a FIFO replay buffer for obs, action, and reward
|
||||
obs_buffer = deque(maxlen=self.buffer_size)
|
||||
next_obs_buffer = deque(maxlen=self.buffer_size)
|
||||
action_buffer = deque(maxlen=self.buffer_size)
|
||||
reward_buffer = deque(maxlen=self.buffer_size)
|
||||
done_buffer = deque(maxlen=self.buffer_size)
|
||||
first_buffer = deque(maxlen=self.buffer_size)
|
||||
|
||||
# Start training loop
|
||||
timer = Timer()
|
||||
run_results = []
|
||||
done_venv = np.zeros((1, self.n_envs))
|
||||
while self.itr < self.n_train_itr:
|
||||
|
||||
# Prepare video paths for each envs --- only applies for the first set of episodes if allowing reset within iteration and each iteration has multiple episodes from one env
|
||||
options_venv = [{} for _ in range(self.n_envs)]
|
||||
if self.itr % self.render_freq == 0 and self.render_video:
|
||||
for env_ind in range(self.n_render):
|
||||
options_venv[env_ind]["video_path"] = os.path.join(
|
||||
self.render_dir, f"itr-{self.itr}_trial-{env_ind}.mp4"
|
||||
)
|
||||
|
||||
# Define train or eval - all envs restart
|
||||
eval_mode = self.itr % self.val_freq == 0 and not self.force_train
|
||||
self.model.eval() if eval_mode else self.model.train()
|
||||
firsts_trajs = np.zeros((self.n_steps + 1, self.n_envs))
|
||||
|
||||
# Reset env at the beginning of an iteration
|
||||
if self.reset_at_iteration or eval_mode or last_itr_eval:
|
||||
prev_obs_venv = self.reset_env_all(options_venv=options_venv)
|
||||
firsts_trajs[0] = 1
|
||||
else:
|
||||
firsts_trajs[0] = (
|
||||
done_venv # if done at the end of last iteration, then the envs are just reset
|
||||
)
|
||||
last_itr_eval = eval_mode
|
||||
reward_trajs = np.empty((0, self.n_envs))
|
||||
|
||||
# Collect a set of trajectories from env
|
||||
for step in range(self.n_steps):
|
||||
if step % 10 == 0:
|
||||
print(f"Processed step {step} of {self.n_steps}")
|
||||
|
||||
# Select action
|
||||
with torch.no_grad():
|
||||
samples = (
|
||||
self.model(
|
||||
cond=torch.from_numpy(prev_obs_venv)
|
||||
.float()
|
||||
.to(self.device),
|
||||
deterministic=eval_mode,
|
||||
)
|
||||
.cpu()
|
||||
.numpy()
|
||||
) # n_env x horizon x act
|
||||
action_venv = samples[:, : self.act_steps]
|
||||
|
||||
# Apply multi-step action
|
||||
obs_venv, reward_venv, done_venv, info_venv = self.venv.step(
|
||||
action_venv
|
||||
)
|
||||
reward_trajs = np.vstack((reward_trajs, reward_venv[None]))
|
||||
|
||||
# add to buffer
|
||||
for i in range(self.n_envs):
|
||||
obs_buffer.append(prev_obs_venv[i])
|
||||
next_obs_buffer.append(obs_venv[i])
|
||||
action_buffer.append(action_venv[i])
|
||||
reward_buffer.append(reward_venv[i] * self.scale_reward_factor)
|
||||
done_buffer.append(done_venv[i])
|
||||
first_buffer.append(firsts_trajs[step])
|
||||
|
||||
firsts_trajs[step + 1] = done_venv
|
||||
prev_obs_venv = obs_venv
|
||||
|
||||
# Summarize episode reward --- this needs to be handled differently depending on whether the environment is reset after each iteration. Only count episodes that finish within the iteration.
|
||||
episodes_start_end = []
|
||||
for env_ind in range(self.n_envs):
|
||||
env_steps = np.where(firsts_trajs[:, env_ind] == 1)[0]
|
||||
for i in range(len(env_steps) - 1):
|
||||
start = env_steps[i]
|
||||
end = env_steps[i + 1]
|
||||
if end - start > 1:
|
||||
episodes_start_end.append((env_ind, start, end - 1))
|
||||
if len(episodes_start_end) > 0:
|
||||
reward_trajs_split = [
|
||||
reward_trajs[start : end + 1, env_ind]
|
||||
for env_ind, start, end in episodes_start_end
|
||||
]
|
||||
num_episode_finished = len(reward_trajs_split)
|
||||
episode_reward = np.array(
|
||||
[np.sum(reward_traj) for reward_traj in reward_trajs_split]
|
||||
)
|
||||
episode_best_reward = np.array(
|
||||
[
|
||||
np.max(reward_traj) / self.act_steps
|
||||
for reward_traj in reward_trajs_split
|
||||
]
|
||||
)
|
||||
avg_episode_reward = np.mean(episode_reward)
|
||||
avg_best_reward = np.mean(episode_best_reward)
|
||||
success_rate = np.mean(
|
||||
episode_best_reward >= self.best_reward_threshold_for_success
|
||||
)
|
||||
else:
|
||||
episode_reward = np.array([])
|
||||
num_episode_finished = 0
|
||||
avg_episode_reward = 0
|
||||
avg_best_reward = 0
|
||||
success_rate = 0
|
||||
log.info("[WARNING] No episode completed within the iteration!")
|
||||
|
||||
if not eval_mode:
|
||||
|
||||
num_batch = self.replay_ratio
|
||||
|
||||
# Critic learning
|
||||
for _ in range(num_batch):
|
||||
# Sample batch
|
||||
inds = np.random.choice(len(obs_buffer), self.batch_size)
|
||||
obs_b = (
|
||||
torch.from_numpy(np.vstack([obs_buffer[i][None] for i in inds]))
|
||||
.float()
|
||||
.to(self.device)
|
||||
)
|
||||
next_obs_b = (
|
||||
torch.from_numpy(
|
||||
np.vstack([next_obs_buffer[i][None] for i in inds])
|
||||
)
|
||||
.float()
|
||||
.to(self.device)
|
||||
)
|
||||
actions_b = (
|
||||
torch.from_numpy(
|
||||
np.vstack([action_buffer[i][None] for i in inds])
|
||||
)
|
||||
.float()
|
||||
.to(self.device)
|
||||
)
|
||||
rewards_b = (
|
||||
torch.from_numpy(np.vstack([reward_buffer[i] for i in inds]))
|
||||
.float()
|
||||
.to(self.device)
|
||||
)
|
||||
dones_b = (
|
||||
torch.from_numpy(np.vstack([done_buffer[i] for i in inds]))
|
||||
.float()
|
||||
.to(self.device)
|
||||
)
|
||||
|
||||
# Update critic
|
||||
loss_critic = self.model.loss_critic(
|
||||
obs_b, next_obs_b, actions_b, rewards_b, dones_b, self.gamma
|
||||
)
|
||||
self.critic_optimizer.zero_grad()
|
||||
loss_critic.backward()
|
||||
self.critic_optimizer.step()
|
||||
|
||||
# Actor learning
|
||||
for _ in range(num_batch):
|
||||
|
||||
# Sample batch
|
||||
inds = np.random.choice(len(obs_buffer), self.batch_size)
|
||||
obs_b = (
|
||||
torch.from_numpy(np.vstack([obs_buffer[i][None] for i in inds]))
|
||||
.float()
|
||||
.to(self.device)
|
||||
)
|
||||
actions_b = (
|
||||
torch.from_numpy(
|
||||
np.vstack([action_buffer[i][None] for i in inds])
|
||||
)
|
||||
.float()
|
||||
.to(self.device)
|
||||
)
|
||||
|
||||
# Replace actions in buffer with guided actions
|
||||
guided_action_list = []
|
||||
|
||||
# get Q-perturbed actions by optimizing
|
||||
actions_flat = actions_b.reshape(actions_b.shape[0], -1)
|
||||
actions_optim = torch.optim.Adam(
|
||||
[actions_flat], lr=self.eta, eps=1e-5
|
||||
)
|
||||
for _ in range(self.action_gradient_steps):
|
||||
actions_flat.requires_grad_(True)
|
||||
q_values_1, q_values_2 = self.model.critic(obs_b, actions_flat)
|
||||
q_values = torch.min(q_values_1, q_values_2)
|
||||
action_opt_loss = -q_values.sum()
|
||||
|
||||
actions_optim.zero_grad()
|
||||
action_opt_loss.backward(torch.ones_like(action_opt_loss))
|
||||
|
||||
# get the perturbed action
|
||||
actions_optim.step()
|
||||
|
||||
actions_flat.requires_grad_(False)
|
||||
actions_flat.clamp_(-1.0, 1.0)
|
||||
guided_action = actions_flat.detach()
|
||||
guided_action = guided_action.reshape(
|
||||
guided_action.shape[0], -1, self.action_dim
|
||||
)
|
||||
guided_action_list.append(guided_action)
|
||||
guided_action_stacked = torch.cat(guided_action_list, 0)
|
||||
|
||||
# Add to buffer (need separate indices since we're working with a limited subset)
|
||||
for i, i_buf in enumerate(inds):
|
||||
action_buffer[i_buf] = (
|
||||
guided_action_stacked[i].detach().cpu().numpy()
|
||||
)
|
||||
|
||||
# Update policy with collected trajectories
|
||||
loss = self.model.loss(guided_action.detach(), {0: obs_b})
|
||||
self.actor_optimizer.zero_grad()
|
||||
loss.backward()
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
if self.max_grad_norm is not None:
|
||||
torch.nn.utils.clip_grad_norm_(
|
||||
self.model.actor.parameters(), self.max_grad_norm
|
||||
)
|
||||
self.actor_optimizer.step()
|
||||
|
||||
# Update lr
|
||||
self.actor_lr_scheduler.step()
|
||||
self.critic_lr_scheduler.step()
|
||||
|
||||
# Save model
|
||||
if self.itr % self.save_model_freq == 0 or self.itr == self.n_train_itr - 1:
|
||||
self.save_model()
|
||||
|
||||
# Log loss and save metrics
|
||||
run_results.append(
|
||||
{
|
||||
"itr": self.itr,
|
||||
}
|
||||
)
|
||||
if self.itr % self.log_freq == 0:
|
||||
if eval_mode:
|
||||
log.info(
|
||||
f"eval: success rate {success_rate:8.4f} | avg episode reward {avg_episode_reward:8.4f} | avg best reward {avg_best_reward:8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"success rate - eval": success_rate,
|
||||
"avg episode reward - eval": avg_episode_reward,
|
||||
"avg best reward - eval": avg_best_reward,
|
||||
"num episode - eval": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=False,
|
||||
)
|
||||
run_results[-1]["eval_success_rate"] = success_rate
|
||||
run_results[-1]["eval_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["eval_best_reward"] = avg_best_reward
|
||||
else:
|
||||
log.info(
|
||||
f"{self.itr}: loss {loss:8.4f} | reward {avg_episode_reward:8.4f} |t:{timer():8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"loss": loss,
|
||||
"loss - critic": loss_critic,
|
||||
"avg episode reward - train": avg_episode_reward,
|
||||
"num episode - train": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=True,
|
||||
)
|
||||
run_results[-1]["loss"] = loss
|
||||
run_results[-1]["loss_critic"] = loss_critic
|
||||
run_results[-1]["train_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["time"] = timer()
|
||||
with open(self.result_path, "wb") as f:
|
||||
pickle.dump(run_results, f)
|
||||
self.itr += 1
|
||||
@@ -0,0 +1,317 @@
|
||||
"""
|
||||
Diffusion Q-Learning (DQL)
|
||||
|
||||
Learns a critic Q-function and backprops the expected Q-value to train the actor
|
||||
|
||||
pi = argmin L_d(\theta) - \alpha * E[Q(s, a)]
|
||||
L_d is demonstration loss for regularization
|
||||
"""
|
||||
|
||||
import os
|
||||
import pickle
|
||||
import numpy as np
|
||||
import torch
|
||||
import logging
|
||||
import wandb
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from util.timer import Timer
|
||||
from collections import deque
|
||||
from agent.finetune.train_agent import TrainAgent
|
||||
from util.scheduler import CosineAnnealingWarmupRestarts
|
||||
|
||||
|
||||
class TrainDQLDiffusionAgent(TrainAgent):
|
||||
|
||||
def __init__(self, cfg):
|
||||
super().__init__(cfg)
|
||||
|
||||
# note the discount factor gamma here is applied to reward every act_steps, instead of every env step
|
||||
self.gamma = cfg.train.gamma
|
||||
|
||||
# Wwarm up period for critic before actor updates
|
||||
self.n_critic_warmup_itr = cfg.train.n_critic_warmup_itr
|
||||
|
||||
# Optimizer
|
||||
self.actor_optimizer = torch.optim.AdamW(
|
||||
self.model.actor.parameters(),
|
||||
lr=cfg.train.actor_lr,
|
||||
weight_decay=cfg.train.actor_weight_decay,
|
||||
)
|
||||
# use cosine scheduler with linear warmup
|
||||
self.actor_lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.actor_optimizer,
|
||||
first_cycle_steps=cfg.train.actor_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.actor_lr,
|
||||
min_lr=cfg.train.actor_lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.actor_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
self.critic_optimizer = torch.optim.AdamW(
|
||||
self.model.critic.parameters(),
|
||||
lr=cfg.train.critic_lr,
|
||||
weight_decay=cfg.train.critic_weight_decay,
|
||||
)
|
||||
self.critic_lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.critic_optimizer,
|
||||
first_cycle_steps=cfg.train.critic_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.critic_lr,
|
||||
min_lr=cfg.train.critic_lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.critic_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
|
||||
# Buffer size
|
||||
self.buffer_size = cfg.train.buffer_size
|
||||
|
||||
# Perturbation scale
|
||||
self.eta = cfg.train.eta
|
||||
|
||||
# Reward factor - scale down mujoco reward for better critic training
|
||||
self.scale_reward_factor = cfg.train.scale_reward_factor
|
||||
|
||||
# Updates
|
||||
self.replay_ratio = cfg.train.replay_ratio
|
||||
|
||||
def run(self):
|
||||
|
||||
# make a FIFO replay buffer for obs, action, and reward
|
||||
obs_buffer = deque(maxlen=self.buffer_size)
|
||||
next_obs_buffer = deque(maxlen=self.buffer_size)
|
||||
action_buffer = deque(maxlen=self.buffer_size)
|
||||
reward_buffer = deque(maxlen=self.buffer_size)
|
||||
done_buffer = deque(maxlen=self.buffer_size)
|
||||
first_buffer = deque(maxlen=self.buffer_size)
|
||||
|
||||
# Start training loop
|
||||
timer = Timer()
|
||||
run_results = []
|
||||
done_venv = np.zeros((1, self.n_envs))
|
||||
while self.itr < self.n_train_itr:
|
||||
|
||||
# Prepare video paths for each envs --- only applies for the first set of episodes if allowing reset within iteration and each iteration has multiple episodes from one env
|
||||
options_venv = [{} for _ in range(self.n_envs)]
|
||||
if self.itr % self.render_freq == 0 and self.render_video:
|
||||
for env_ind in range(self.n_render):
|
||||
options_venv[env_ind]["video_path"] = os.path.join(
|
||||
self.render_dir, f"itr-{self.itr}_trial-{env_ind}.mp4"
|
||||
)
|
||||
|
||||
# Define train or eval - all envs restart
|
||||
eval_mode = self.itr % self.val_freq == 0 and not self.force_train
|
||||
self.model.eval() if eval_mode else self.model.train()
|
||||
firsts_trajs = np.zeros((self.n_steps + 1, self.n_envs))
|
||||
|
||||
# Reset env at the beginning of an iteration
|
||||
if self.reset_at_iteration or eval_mode or last_itr_eval:
|
||||
prev_obs_venv = self.reset_env_all(options_venv=options_venv)
|
||||
firsts_trajs[0] = 1
|
||||
else:
|
||||
firsts_trajs[0] = (
|
||||
done_venv # if done at the end of last iteration, then the envs are just reset
|
||||
)
|
||||
last_itr_eval = eval_mode
|
||||
reward_trajs = np.empty((0, self.n_envs))
|
||||
|
||||
# Collect a set of trajectories from env
|
||||
for step in range(self.n_steps):
|
||||
if step % 10 == 0:
|
||||
print(f"Processed step {step} of {self.n_steps}")
|
||||
|
||||
# Select action
|
||||
with torch.no_grad():
|
||||
samples = (
|
||||
self.model(
|
||||
cond=torch.from_numpy(prev_obs_venv)
|
||||
.float()
|
||||
.to(self.device),
|
||||
deterministic=eval_mode,
|
||||
)
|
||||
.cpu()
|
||||
.numpy()
|
||||
) # n_env x horizon x act
|
||||
action_venv = samples[:, : self.act_steps]
|
||||
|
||||
# Apply multi-step action
|
||||
obs_venv, reward_venv, done_venv, info_venv = self.venv.step(
|
||||
action_venv
|
||||
)
|
||||
reward_trajs = np.vstack((reward_trajs, reward_venv[None]))
|
||||
|
||||
# add to buffer
|
||||
for i in range(self.n_envs):
|
||||
obs_buffer.append(prev_obs_venv[i])
|
||||
next_obs_buffer.append(obs_venv[i])
|
||||
action_buffer.append(action_venv[i])
|
||||
reward_buffer.append(reward_venv[i] * self.scale_reward_factor)
|
||||
done_buffer.append(done_venv[i])
|
||||
first_buffer.append(firsts_trajs[step])
|
||||
|
||||
firsts_trajs[step + 1] = done_venv
|
||||
prev_obs_venv = obs_venv
|
||||
|
||||
# Summarize episode reward --- this needs to be handled differently depending on whether the environment is reset after each iteration. Only count episodes that finish within the iteration.
|
||||
episodes_start_end = []
|
||||
for env_ind in range(self.n_envs):
|
||||
env_steps = np.where(firsts_trajs[:, env_ind] == 1)[0]
|
||||
for i in range(len(env_steps) - 1):
|
||||
start = env_steps[i]
|
||||
end = env_steps[i + 1]
|
||||
if end - start > 1:
|
||||
episodes_start_end.append((env_ind, start, end - 1))
|
||||
if len(episodes_start_end) > 0:
|
||||
reward_trajs_split = [
|
||||
reward_trajs[start : end + 1, env_ind]
|
||||
for env_ind, start, end in episodes_start_end
|
||||
]
|
||||
num_episode_finished = len(reward_trajs_split)
|
||||
episode_reward = np.array(
|
||||
[np.sum(reward_traj) for reward_traj in reward_trajs_split]
|
||||
)
|
||||
episode_best_reward = np.array(
|
||||
[
|
||||
np.max(reward_traj) / self.act_steps
|
||||
for reward_traj in reward_trajs_split
|
||||
]
|
||||
)
|
||||
avg_episode_reward = np.mean(episode_reward)
|
||||
avg_best_reward = np.mean(episode_best_reward)
|
||||
success_rate = np.mean(
|
||||
episode_best_reward >= self.best_reward_threshold_for_success
|
||||
)
|
||||
else:
|
||||
episode_reward = np.array([])
|
||||
num_episode_finished = 0
|
||||
avg_episode_reward = 0
|
||||
avg_best_reward = 0
|
||||
success_rate = 0
|
||||
log.info("[WARNING] No episode completed within the iteration!")
|
||||
|
||||
if not eval_mode:
|
||||
|
||||
num_batch = self.replay_ratio
|
||||
|
||||
# Critic learning
|
||||
for _ in range(num_batch):
|
||||
# Sample batch
|
||||
inds = np.random.choice(len(obs_buffer), self.batch_size)
|
||||
obs_b = (
|
||||
torch.from_numpy(np.vstack([obs_buffer[i][None] for i in inds]))
|
||||
.float()
|
||||
.to(self.device)
|
||||
)
|
||||
next_obs_b = (
|
||||
torch.from_numpy(
|
||||
np.vstack([next_obs_buffer[i][None] for i in inds])
|
||||
)
|
||||
.float()
|
||||
.to(self.device)
|
||||
)
|
||||
actions_b = (
|
||||
torch.from_numpy(
|
||||
np.vstack([action_buffer[i][None] for i in inds])
|
||||
)
|
||||
.float()
|
||||
.to(self.device)
|
||||
)
|
||||
rewards_b = (
|
||||
torch.from_numpy(np.vstack([reward_buffer[i] for i in inds]))
|
||||
.float()
|
||||
.to(self.device)
|
||||
)
|
||||
dones_b = (
|
||||
torch.from_numpy(np.vstack([done_buffer[i] for i in inds]))
|
||||
.float()
|
||||
.to(self.device)
|
||||
)
|
||||
|
||||
# Update critic
|
||||
loss_critic = self.model.loss_critic(
|
||||
obs_b, next_obs_b, actions_b, rewards_b, dones_b, self.gamma
|
||||
)
|
||||
self.critic_optimizer.zero_grad()
|
||||
loss_critic.backward()
|
||||
self.critic_optimizer.step()
|
||||
|
||||
# get the new action and q values
|
||||
samples = self.model.forward_train(
|
||||
cond=obs_b.to(self.device),
|
||||
deterministic=eval_mode,
|
||||
)
|
||||
output_venv = samples # n_env x horizon x act
|
||||
action_venv = output_venv[:, : self.act_steps, : self.action_dim]
|
||||
actions_flat_b = action_venv.reshape(action_venv.shape[0], -1)
|
||||
q_values_b = self.model.critic(obs_b, actions_flat_b)
|
||||
q1_new_action, q2_new_action = q_values_b
|
||||
|
||||
# Update policy with collected trajectories
|
||||
self.actor_optimizer.zero_grad()
|
||||
actor_loss = self.model.loss_actor(
|
||||
obs_b, actions_b, q1_new_action, q2_new_action, self.eta
|
||||
)
|
||||
actor_loss.backward()
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
if self.max_grad_norm is not None:
|
||||
torch.nn.utils.clip_grad_norm_(
|
||||
self.model.actor.parameters(), self.max_grad_norm
|
||||
)
|
||||
self.actor_optimizer.step()
|
||||
loss = actor_loss
|
||||
|
||||
# Update lr
|
||||
self.actor_lr_scheduler.step()
|
||||
self.critic_lr_scheduler.step()
|
||||
|
||||
# Save model
|
||||
if self.itr % self.save_model_freq == 0 or self.itr == self.n_train_itr - 1:
|
||||
self.save_model()
|
||||
|
||||
# Log loss and save metrics
|
||||
run_results.append(
|
||||
{
|
||||
"itr": self.itr,
|
||||
}
|
||||
)
|
||||
if self.itr % self.log_freq == 0:
|
||||
if eval_mode:
|
||||
log.info(
|
||||
f"eval: success rate {success_rate:8.4f} | avg episode reward {avg_episode_reward:8.4f} | avg best reward {avg_best_reward:8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"success rate - eval": success_rate,
|
||||
"avg episode reward - eval": avg_episode_reward,
|
||||
"avg best reward - eval": avg_best_reward,
|
||||
"num episode - eval": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=False,
|
||||
)
|
||||
run_results[-1]["eval_success_rate"] = success_rate
|
||||
run_results[-1]["eval_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["eval_best_reward"] = avg_best_reward
|
||||
else:
|
||||
log.info(
|
||||
f"{self.itr}: loss {loss:8.4f} | reward {avg_episode_reward:8.4f} |t:{timer():8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"loss": loss,
|
||||
"loss - critic": loss_critic,
|
||||
"avg episode reward - train": avg_episode_reward,
|
||||
"num episode - train": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=True,
|
||||
)
|
||||
run_results[-1]["loss"] = loss
|
||||
run_results[-1]["loss_critic"] = loss_critic
|
||||
run_results[-1]["train_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["time"] = timer()
|
||||
with open(self.result_path, "wb") as f:
|
||||
pickle.dump(run_results, f)
|
||||
self.itr += 1
|
||||
@@ -0,0 +1,347 @@
|
||||
"""
|
||||
Implicit diffusion Q-learning (IDQL) trainer for diffusion policy.
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
import pickle
|
||||
import einops
|
||||
import numpy as np
|
||||
import torch
|
||||
import logging
|
||||
import wandb
|
||||
from copy import deepcopy
|
||||
import random
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from util.timer import Timer
|
||||
from collections import deque
|
||||
from agent.finetune.train_agent import TrainAgent
|
||||
from util.scheduler import CosineAnnealingWarmupRestarts
|
||||
|
||||
|
||||
class TrainIDQLDiffusionAgent(TrainAgent):
|
||||
|
||||
def __init__(self, cfg):
|
||||
super().__init__(cfg)
|
||||
|
||||
# note the discount factor gamma here is applied to reward every act_steps, instead of every env step
|
||||
self.gamma = cfg.train.gamma
|
||||
|
||||
# Wwarm up period for critic before actor updates
|
||||
self.n_critic_warmup_itr = cfg.train.n_critic_warmup_itr
|
||||
|
||||
# Optimizer
|
||||
self.actor_optimizer = torch.optim.AdamW(
|
||||
self.model.actor.parameters(),
|
||||
lr=cfg.train.actor_lr,
|
||||
weight_decay=cfg.train.actor_weight_decay,
|
||||
)
|
||||
self.actor_lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.actor_optimizer,
|
||||
first_cycle_steps=cfg.train.actor_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.actor_lr,
|
||||
min_lr=cfg.train.actor_lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.actor_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
self.critic_q_optimizer = torch.optim.AdamW(
|
||||
self.model.critic_q.parameters(),
|
||||
lr=cfg.train.critic_lr,
|
||||
weight_decay=cfg.train.critic_weight_decay,
|
||||
)
|
||||
self.critic_v_optimizer = torch.optim.AdamW(
|
||||
self.model.critic_v.parameters(),
|
||||
lr=cfg.train.critic_lr,
|
||||
weight_decay=cfg.train.critic_weight_decay,
|
||||
)
|
||||
self.critic_v_lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.critic_v_optimizer,
|
||||
first_cycle_steps=cfg.train.critic_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.critic_lr,
|
||||
min_lr=cfg.train.critic_lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.critic_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
self.critic_q_lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.critic_q_optimizer,
|
||||
first_cycle_steps=cfg.train.critic_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.critic_lr,
|
||||
min_lr=cfg.train.critic_lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.critic_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
|
||||
# Buffer size
|
||||
self.buffer_size = cfg.train.buffer_size
|
||||
|
||||
# Actor params
|
||||
self.use_expectile_exploration = cfg.train.use_expectile_exploration
|
||||
|
||||
# Scaling reward
|
||||
self.scale_reward_factor = cfg.train.scale_reward_factor
|
||||
|
||||
# Updates
|
||||
self.replay_ratio = cfg.train.replay_ratio
|
||||
self.critic_tau = cfg.train.critic_tau
|
||||
|
||||
# Whether to use deterministic mode when sampling at eval
|
||||
self.eval_deterministic = cfg.train.get("eval_deterministic", False)
|
||||
|
||||
# Sampling
|
||||
self.num_sample = cfg.train.eval_sample_num
|
||||
|
||||
def run(self):
|
||||
|
||||
# make a FIFO replay buffer for obs, action, and reward
|
||||
obs_buffer = deque(maxlen=self.buffer_size)
|
||||
action_buffer = deque(maxlen=self.buffer_size)
|
||||
next_obs_buffer = deque(maxlen=self.buffer_size)
|
||||
reward_buffer = deque(maxlen=self.buffer_size)
|
||||
done_buffer = deque(maxlen=self.buffer_size)
|
||||
first_buffer = deque(maxlen=self.buffer_size)
|
||||
|
||||
# Start training loop
|
||||
timer = Timer()
|
||||
run_results = []
|
||||
done_venv = np.zeros((1, self.n_envs))
|
||||
while self.itr < self.n_train_itr:
|
||||
|
||||
# Prepare video paths for each envs --- only applies for the first set of episodes if allowing reset within iteration and each iteration has multiple episodes from one env
|
||||
options_venv = [{} for _ in range(self.n_envs)]
|
||||
if self.itr % self.render_freq == 0 and self.render_video:
|
||||
for env_ind in range(self.n_render):
|
||||
options_venv[env_ind]["video_path"] = os.path.join(
|
||||
self.render_dir, f"itr-{self.itr}_trial-{env_ind}.mp4"
|
||||
)
|
||||
|
||||
# Define train or eval - all envs restart
|
||||
eval_mode = self.itr % self.val_freq == 0 and not self.force_train
|
||||
self.model.eval() if eval_mode else self.model.train()
|
||||
firsts_trajs = np.zeros((self.n_steps + 1, self.n_envs))
|
||||
|
||||
# Reset env at the beginning of an iteration
|
||||
if self.reset_at_iteration or eval_mode or last_itr_eval:
|
||||
prev_obs_venv = self.reset_env_all(options_venv=options_venv)
|
||||
firsts_trajs[0] = 1
|
||||
else:
|
||||
firsts_trajs[0] = (
|
||||
done_venv # if done at the end of last iteration, then the envs are just reset
|
||||
)
|
||||
last_itr_eval = eval_mode
|
||||
reward_trajs = np.empty((0, self.n_envs))
|
||||
|
||||
# Collect a set of trajectories from env
|
||||
for step in range(self.n_steps):
|
||||
if step % 10 == 0:
|
||||
print(f"Processed step {step} of {self.n_steps}")
|
||||
|
||||
# Select action
|
||||
with torch.no_grad():
|
||||
samples = (
|
||||
self.model(
|
||||
cond=torch.from_numpy(prev_obs_venv)
|
||||
.float()
|
||||
.to(self.device),
|
||||
deterministic=eval_mode and self.eval_deterministic,
|
||||
num_sample=self.num_sample,
|
||||
use_expectile_exploration=self.use_expectile_exploration,
|
||||
)
|
||||
.cpu()
|
||||
.numpy()
|
||||
) # n_env x horizon x act
|
||||
action_venv = samples[:, : self.act_steps]
|
||||
|
||||
# Apply multi-step action
|
||||
obs_venv, reward_venv, done_venv, info_venv = self.venv.step(
|
||||
action_venv
|
||||
)
|
||||
reward_trajs = np.vstack((reward_trajs, reward_venv[None]))
|
||||
|
||||
# add to buffer
|
||||
obs_buffer.append(prev_obs_venv)
|
||||
action_buffer.append(action_venv)
|
||||
next_obs_buffer.append(obs_venv)
|
||||
reward_buffer.append(reward_venv * self.scale_reward_factor)
|
||||
done_buffer.append(done_venv)
|
||||
first_buffer.append(firsts_trajs[step])
|
||||
|
||||
firsts_trajs[step + 1] = done_venv
|
||||
prev_obs_venv = obs_venv
|
||||
|
||||
# Summarize episode reward --- this needs to be handled differently depending on whether the environment is reset after each iteration. Only count episodes that finish within the iteration.
|
||||
episodes_start_end = []
|
||||
for env_ind in range(self.n_envs):
|
||||
env_steps = np.where(firsts_trajs[:, env_ind] == 1)[0]
|
||||
for i in range(len(env_steps) - 1):
|
||||
start = env_steps[i]
|
||||
end = env_steps[i + 1]
|
||||
if end - start > 1:
|
||||
episodes_start_end.append((env_ind, start, end - 1))
|
||||
if len(episodes_start_end) > 0:
|
||||
reward_trajs_split = [
|
||||
reward_trajs[start : end + 1, env_ind]
|
||||
for env_ind, start, end in episodes_start_end
|
||||
]
|
||||
num_episode_finished = len(reward_trajs_split)
|
||||
episode_reward = np.array(
|
||||
[np.sum(reward_traj) for reward_traj in reward_trajs_split]
|
||||
)
|
||||
episode_best_reward = np.array(
|
||||
[
|
||||
np.max(reward_traj) / self.act_steps
|
||||
for reward_traj in reward_trajs_split
|
||||
]
|
||||
)
|
||||
avg_episode_reward = np.mean(episode_reward)
|
||||
avg_best_reward = np.mean(episode_best_reward)
|
||||
success_rate = np.mean(
|
||||
episode_best_reward >= self.best_reward_threshold_for_success
|
||||
)
|
||||
else:
|
||||
episode_reward = np.array([])
|
||||
num_episode_finished = 0
|
||||
avg_episode_reward = 0
|
||||
avg_best_reward = 0
|
||||
success_rate = 0
|
||||
log.info("[WARNING] No episode completed within the iteration!")
|
||||
|
||||
# Update
|
||||
if not eval_mode:
|
||||
|
||||
obs_trajs = np.array(deepcopy(obs_buffer))
|
||||
action_trajs = np.array(deepcopy(action_buffer))
|
||||
next_obs_trajs = np.array(deepcopy(next_obs_buffer))
|
||||
reward_trajs = np.array(deepcopy(reward_buffer))
|
||||
done_trajs = np.array(deepcopy(done_buffer))
|
||||
first_trajs = np.array(deepcopy(first_buffer))
|
||||
|
||||
# flatten
|
||||
obs_trajs = einops.rearrange(
|
||||
obs_trajs,
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
next_obs_trajs = einops.rearrange(
|
||||
next_obs_trajs,
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
action_trajs = einops.rearrange(
|
||||
action_trajs,
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
reward_trajs = reward_trajs.reshape(-1)
|
||||
done_trajs = done_trajs.reshape(-1)
|
||||
first_trajs = first_trajs.reshape(-1)
|
||||
|
||||
num_batch = int(
|
||||
self.n_steps * self.n_envs / self.batch_size * self.replay_ratio
|
||||
)
|
||||
|
||||
for _ in range(num_batch):
|
||||
|
||||
# Sample batch
|
||||
inds = np.random.choice(len(obs_trajs), self.batch_size)
|
||||
obs_b = torch.from_numpy(obs_trajs[inds]).float().to(self.device)
|
||||
next_obs_b = (
|
||||
torch.from_numpy(next_obs_trajs[inds]).float().to(self.device)
|
||||
)
|
||||
actions_b = (
|
||||
torch.from_numpy(action_trajs[inds]).float().to(self.device)
|
||||
)
|
||||
reward_b = (
|
||||
torch.from_numpy(reward_trajs[inds]).float().to(self.device)
|
||||
)
|
||||
done_b = torch.from_numpy(done_trajs[inds]).float().to(self.device)
|
||||
|
||||
# update critic value function
|
||||
critic_loss_v = self.model.loss_critic_v(obs_b, actions_b)
|
||||
self.critic_v_optimizer.zero_grad()
|
||||
critic_loss_v.backward()
|
||||
self.critic_v_optimizer.step()
|
||||
|
||||
# update critic q function
|
||||
critic_loss_q = self.model.loss_critic_q(
|
||||
obs_b, next_obs_b, actions_b, reward_b, done_b, self.gamma
|
||||
)
|
||||
self.critic_q_optimizer.zero_grad()
|
||||
critic_loss_q.backward()
|
||||
self.critic_q_optimizer.step()
|
||||
|
||||
# update target q function
|
||||
self.model.update_target_critic(self.critic_tau)
|
||||
|
||||
loss_critic = critic_loss_q.detach() + critic_loss_v.detach()
|
||||
|
||||
# Update policy with collected trajectories - no weighting
|
||||
loss = self.model.loss(
|
||||
actions_b,
|
||||
obs_b,
|
||||
)
|
||||
self.actor_optimizer.zero_grad()
|
||||
loss.backward()
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
if self.max_grad_norm is not None:
|
||||
torch.nn.utils.clip_grad_norm_(
|
||||
self.model.actor.parameters(), self.max_grad_norm
|
||||
)
|
||||
self.actor_optimizer.step()
|
||||
|
||||
# Update lr
|
||||
self.actor_lr_scheduler.step()
|
||||
self.critic_v_lr_scheduler.step()
|
||||
self.critic_q_lr_scheduler.step()
|
||||
|
||||
# Save model
|
||||
if self.itr % self.save_model_freq == 0 or self.itr == self.n_train_itr - 1:
|
||||
self.save_model()
|
||||
|
||||
# Log loss and save metrics
|
||||
run_results.append(
|
||||
{
|
||||
"itr": self.itr,
|
||||
}
|
||||
)
|
||||
if self.itr % self.log_freq == 0:
|
||||
if eval_mode:
|
||||
log.info(
|
||||
f"eval: success rate {success_rate:8.4f} | avg episode reward {avg_episode_reward:8.4f} | avg best reward {avg_best_reward:8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"success rate - eval": success_rate,
|
||||
"avg episode reward - eval": avg_episode_reward,
|
||||
"avg best reward - eval": avg_best_reward,
|
||||
"num episode - eval": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=False,
|
||||
)
|
||||
run_results[-1]["eval_success_rate"] = success_rate
|
||||
run_results[-1]["eval_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["eval_best_reward"] = avg_best_reward
|
||||
else:
|
||||
log.info(
|
||||
f"{self.itr}: loss {loss:8.4f} | reward {avg_episode_reward:8.4f} |t:{timer():8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"loss": loss,
|
||||
"loss - critic": loss_critic,
|
||||
"avg episode reward - train": avg_episode_reward,
|
||||
"num episode - train": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=True,
|
||||
)
|
||||
run_results[-1]["loss"] = loss
|
||||
run_results[-1]["loss_critic"] = loss_critic
|
||||
run_results[-1]["train_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["time"] = timer()
|
||||
with open(self.result_path, "wb") as f:
|
||||
pickle.dump(run_results, f)
|
||||
self.itr += 1
|
||||
@@ -0,0 +1,112 @@
|
||||
"""
|
||||
Parent PPO fine-tuning agent class.
|
||||
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
import torch
|
||||
import logging
|
||||
from util.scheduler import CosineAnnealingWarmupRestarts
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from agent.finetune.train_agent import TrainAgent
|
||||
from util.reward_scaling import RunningRewardScaler
|
||||
|
||||
|
||||
class TrainPPOAgent(TrainAgent):
|
||||
|
||||
def __init__(self, cfg):
|
||||
super().__init__(cfg)
|
||||
|
||||
# Batch size for logprobs calculations after an iteration --- prevent out of memory if using a single batch
|
||||
self.logprob_batch_size = cfg.train.get("logprob_batch_size", 10000)
|
||||
assert (
|
||||
self.logprob_batch_size % self.n_envs == 0
|
||||
), "logprob_batch_size must be divisible by n_envs"
|
||||
|
||||
# note the discount factor gamma here is applied to reward every act_steps, instead of every env step
|
||||
self.gamma = cfg.train.gamma
|
||||
|
||||
# Wwarm up period for critic before actor updates
|
||||
self.n_critic_warmup_itr = cfg.train.n_critic_warmup_itr
|
||||
|
||||
# Optimizer
|
||||
self.actor_optimizer = torch.optim.AdamW(
|
||||
self.model.actor_ft.parameters(),
|
||||
lr=cfg.train.actor_lr,
|
||||
weight_decay=cfg.train.actor_weight_decay,
|
||||
)
|
||||
# use cosine scheduler with linear warmup
|
||||
self.actor_lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.actor_optimizer,
|
||||
first_cycle_steps=cfg.train.actor_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.actor_lr,
|
||||
min_lr=cfg.train.actor_lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.actor_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
self.critic_optimizer = torch.optim.AdamW(
|
||||
self.model.critic.parameters(),
|
||||
lr=cfg.train.critic_lr,
|
||||
weight_decay=cfg.train.critic_weight_decay,
|
||||
)
|
||||
self.critic_lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.critic_optimizer,
|
||||
first_cycle_steps=cfg.train.critic_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.critic_lr,
|
||||
min_lr=cfg.train.critic_lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.critic_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
|
||||
# Generalized advantage estimation
|
||||
self.gae_lambda: float = cfg.train.get("gae_lambda", 0.95)
|
||||
|
||||
# If specified, stop gradient update once KL difference reaches it
|
||||
self.target_kl: Optional[float] = cfg.train.target_kl
|
||||
|
||||
# Number of times the collected data is used in gradient update
|
||||
self.update_epochs: int = cfg.train.update_epochs
|
||||
|
||||
# Entropy loss coefficient
|
||||
self.ent_coef: float = cfg.train.get("ent_coef", 0)
|
||||
|
||||
# Value loss coefficient
|
||||
self.vf_coef: float = cfg.train.get("vf_coef", 0)
|
||||
|
||||
# Whether to use running reward scaling
|
||||
self.reward_scale_running: bool = cfg.train.reward_scale_running
|
||||
if self.reward_scale_running:
|
||||
self.running_reward_scaler = RunningRewardScaler(self.n_envs)
|
||||
|
||||
# Scaling reward with constant
|
||||
self.reward_scale_const: float = cfg.train.get("reward_scale_const", 1)
|
||||
|
||||
# Use base policy
|
||||
self.use_bc_loss: bool = cfg.train.get("use_bc_loss", False)
|
||||
self.bc_loss_coeff: float = cfg.train.get("bc_loss_coeff", 0)
|
||||
|
||||
def reset_actor_optimizer(self):
|
||||
"""Not used anywhere currently"""
|
||||
new_optimizer = torch.optim.AdamW(
|
||||
self.model.actor_ft.parameters(),
|
||||
lr=self.cfg.train.actor_lr,
|
||||
weight_decay=self.cfg.train.actor_weight_decay,
|
||||
)
|
||||
new_optimizer.load_state_dict(self.actor_optimizer.state_dict())
|
||||
self.actor_optimizer = new_optimizer
|
||||
|
||||
new_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.actor_optimizer,
|
||||
first_cycle_steps=self.cfg.train.actor_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=self.cfg.train.actor_lr,
|
||||
min_lr=self.cfg.train.actor_lr_scheduler.min_lr,
|
||||
warmup_steps=self.cfg.train.actor_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
new_scheduler.load_state_dict(self.actor_lr_scheduler.state_dict())
|
||||
self.actor_lr_scheduler = new_scheduler
|
||||
log.info("Reset actor optimizer")
|
||||
@@ -0,0 +1,452 @@
|
||||
"""
|
||||
DPPO fine-tuning.
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
import pickle
|
||||
import einops
|
||||
import numpy as np
|
||||
import torch
|
||||
import logging
|
||||
import wandb
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from util.timer import Timer
|
||||
from agent.finetune.train_ppo_agent import TrainPPOAgent
|
||||
from util.scheduler import CosineAnnealingWarmupRestarts
|
||||
|
||||
|
||||
class TrainPPODiffusionAgent(TrainPPOAgent):
|
||||
|
||||
def __init__(self, cfg):
|
||||
super().__init__(cfg)
|
||||
|
||||
# Reward horizon --- always set to act_steps for now
|
||||
self.reward_horizon = cfg.get("reward_horizon", self.act_steps)
|
||||
|
||||
# Eta - between DDIM (=0 for eval) and DDPM (=1 for training)
|
||||
self.learn_eta = self.model.learn_eta
|
||||
if self.learn_eta:
|
||||
self.eta_update_interval = cfg.train.eta_update_interval
|
||||
self.eta_optimizer = torch.optim.AdamW(
|
||||
self.model.eta.parameters(),
|
||||
lr=cfg.train.eta_lr,
|
||||
weight_decay=cfg.train.eta_weight_decay,
|
||||
)
|
||||
self.eta_lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.eta_optimizer,
|
||||
first_cycle_steps=cfg.train.eta_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.eta_lr,
|
||||
min_lr=cfg.train.eta_lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.eta_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
|
||||
def run(self):
|
||||
|
||||
# Start training loop
|
||||
timer = Timer()
|
||||
run_results = []
|
||||
last_itr_eval = False
|
||||
done_venv = np.zeros((1, self.n_envs))
|
||||
while self.itr < self.n_train_itr:
|
||||
|
||||
# Prepare video paths for each envs --- only applies for the first set of episodes if allowing reset within iteration and each iteration has multiple episodes from one env
|
||||
options_venv = [{} for _ in range(self.n_envs)]
|
||||
if self.itr % self.render_freq == 0 and self.render_video:
|
||||
for env_ind in range(self.n_render):
|
||||
options_venv[env_ind]["video_path"] = os.path.join(
|
||||
self.render_dir, f"itr-{self.itr}_trial-{env_ind}.mp4"
|
||||
)
|
||||
|
||||
# Define train or eval - all envs restart
|
||||
eval_mode = self.itr % self.val_freq == 0 and not self.force_train
|
||||
self.model.eval() if eval_mode else self.model.train()
|
||||
last_itr_eval = eval_mode
|
||||
|
||||
# Reset env before iteration starts (1) if specified, (2) at eval mode, or (3) right after eval mode
|
||||
dones_trajs = np.empty((0, self.n_envs))
|
||||
firsts_trajs = np.zeros((self.n_steps + 1, self.n_envs))
|
||||
if self.reset_at_iteration or eval_mode or last_itr_eval:
|
||||
prev_obs_venv = self.reset_env_all(options_venv=options_venv)
|
||||
firsts_trajs[0] = 1
|
||||
else:
|
||||
firsts_trajs[0] = (
|
||||
done_venv # if done at the end of last iteration, then the envs are just reset
|
||||
)
|
||||
|
||||
# Holder
|
||||
obs_trajs = np.empty((0, self.n_envs, self.n_cond_step, self.obs_dim))
|
||||
chains_trajs = np.empty(
|
||||
(
|
||||
0,
|
||||
self.n_envs,
|
||||
self.model.ft_denoising_steps + 1,
|
||||
self.horizon_steps,
|
||||
self.action_dim,
|
||||
)
|
||||
)
|
||||
reward_trajs = np.empty((0, self.n_envs))
|
||||
obs_full_trajs = np.empty((0, self.n_envs, self.obs_dim))
|
||||
obs_full_trajs = np.vstack(
|
||||
(obs_full_trajs, prev_obs_venv[None].squeeze(2))
|
||||
) # remove cond_step dim
|
||||
|
||||
# Collect a set of trajectories from env
|
||||
for step in range(self.n_steps):
|
||||
if step % 10 == 0:
|
||||
print(f"Processed step {step} of {self.n_steps}")
|
||||
|
||||
# Select action
|
||||
with torch.no_grad():
|
||||
samples = self.model(
|
||||
cond=torch.from_numpy(prev_obs_venv).float().to(self.device),
|
||||
deterministic=eval_mode,
|
||||
return_chain=True,
|
||||
)
|
||||
output_venv = (
|
||||
samples.trajectories.cpu().numpy()
|
||||
) # n_env x horizon x act
|
||||
chains_venv = (
|
||||
samples.chains.cpu().numpy()
|
||||
) # n_env x denoising x horizon x act
|
||||
action_venv = output_venv[:, : self.act_steps]
|
||||
|
||||
# Apply multi-step action
|
||||
obs_venv, reward_venv, done_venv, info_venv = self.venv.step(
|
||||
action_venv
|
||||
)
|
||||
if self.save_full_observations:
|
||||
obs_full_venv = np.vstack(
|
||||
[info["full_obs"][None] for info in info_venv]
|
||||
) # n_envs x n_act_steps x obs_dim
|
||||
obs_full_trajs = np.vstack(
|
||||
(obs_full_trajs, obs_full_venv.transpose(1, 0, 2))
|
||||
)
|
||||
obs_trajs = np.vstack((obs_trajs, prev_obs_venv[None]))
|
||||
chains_trajs = np.vstack((chains_trajs, chains_venv[None]))
|
||||
reward_trajs = np.vstack((reward_trajs, reward_venv[None]))
|
||||
dones_trajs = np.vstack((dones_trajs, done_venv[None]))
|
||||
firsts_trajs[step + 1] = done_venv
|
||||
prev_obs_venv = obs_venv
|
||||
|
||||
# Summarize episode reward --- this needs to be handled differently depending on whether the environment is reset after each iteration. Only count episodes that finish within the iteration.
|
||||
episodes_start_end = []
|
||||
for env_ind in range(self.n_envs):
|
||||
env_steps = np.where(firsts_trajs[:, env_ind] == 1)[0]
|
||||
for i in range(len(env_steps) - 1):
|
||||
start = env_steps[i]
|
||||
end = env_steps[i + 1]
|
||||
if end - start > 1:
|
||||
episodes_start_end.append((env_ind, start, end - 1))
|
||||
if len(episodes_start_end) > 0:
|
||||
reward_trajs_split = [
|
||||
reward_trajs[start : end + 1, env_ind]
|
||||
for env_ind, start, end in episodes_start_end
|
||||
]
|
||||
num_episode_finished = len(reward_trajs_split)
|
||||
episode_reward = np.array(
|
||||
[np.sum(reward_traj) for reward_traj in reward_trajs_split]
|
||||
)
|
||||
if (
|
||||
self.furniture_sparse_reward
|
||||
): # only for furniture tasks, where reward only occurs in one env step
|
||||
episode_best_reward = episode_reward
|
||||
else:
|
||||
episode_best_reward = np.array(
|
||||
[
|
||||
np.max(reward_traj) / self.act_steps
|
||||
for reward_traj in reward_trajs_split
|
||||
]
|
||||
)
|
||||
avg_episode_reward = np.mean(episode_reward)
|
||||
avg_best_reward = np.mean(episode_best_reward)
|
||||
success_rate = np.mean(
|
||||
episode_best_reward >= self.best_reward_threshold_for_success
|
||||
)
|
||||
else:
|
||||
episode_reward = np.array([])
|
||||
num_episode_finished = 0
|
||||
avg_episode_reward = 0
|
||||
avg_best_reward = 0
|
||||
success_rate = 0
|
||||
log.info("[WARNING] No episode completed within the iteration!")
|
||||
|
||||
# Update models
|
||||
if not eval_mode:
|
||||
with torch.no_grad():
|
||||
# Calculate value and logprobs - split into batches to prevent out of memory
|
||||
obs_t = einops.rearrange(
|
||||
torch.from_numpy(obs_trajs).float().to(self.device),
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
obs_ts = torch.split(obs_t, self.logprob_batch_size, dim=0)
|
||||
values_trajs = np.empty((0, self.n_envs))
|
||||
for obs in obs_ts:
|
||||
values = self.model.critic(obs).cpu().numpy().flatten()
|
||||
values_trajs = np.vstack(
|
||||
(values_trajs, values.reshape(-1, self.n_envs))
|
||||
)
|
||||
chains_t = einops.rearrange(
|
||||
torch.from_numpy(chains_trajs).float().to(self.device),
|
||||
"s e t h d -> (s e) t h d",
|
||||
)
|
||||
chains_ts = torch.split(chains_t, self.logprob_batch_size, dim=0)
|
||||
logprobs_trajs = np.empty(
|
||||
(
|
||||
0,
|
||||
self.model.ft_denoising_steps,
|
||||
self.horizon_steps,
|
||||
self.action_dim,
|
||||
)
|
||||
)
|
||||
for obs, chains in zip(obs_ts, chains_ts):
|
||||
logprobs = self.model.get_logprobs(obs, chains).cpu().numpy()
|
||||
logprobs_trajs = np.vstack(
|
||||
(
|
||||
logprobs_trajs,
|
||||
logprobs.reshape(-1, *logprobs_trajs.shape[1:]),
|
||||
)
|
||||
)
|
||||
|
||||
# normalize reward with running variance if specified
|
||||
if self.reward_scale_running:
|
||||
reward_trajs_transpose = self.running_reward_scaler(
|
||||
reward=reward_trajs.T, first=firsts_trajs[:-1].T
|
||||
)
|
||||
reward_trajs = reward_trajs_transpose.T
|
||||
|
||||
# bootstrap value with GAE if not done - apply reward scaling with constant if specified
|
||||
obs_venv_ts = torch.from_numpy(obs_venv).float().to(self.device)
|
||||
with torch.no_grad():
|
||||
next_value = (
|
||||
self.model.critic(obs_venv_ts).reshape(1, -1).cpu().numpy()
|
||||
)
|
||||
advantages_trajs = np.zeros_like(reward_trajs)
|
||||
lastgaelam = 0
|
||||
for t in reversed(range(self.n_steps)):
|
||||
if t == self.n_steps - 1:
|
||||
nextnonterminal = 1.0 - done_venv
|
||||
nextvalues = next_value
|
||||
else:
|
||||
nextnonterminal = 1.0 - dones_trajs[t + 1]
|
||||
nextvalues = values_trajs[t + 1]
|
||||
# delta = r + gamma*V(st+1) - V(st)
|
||||
delta = (
|
||||
reward_trajs[t] * self.reward_scale_const
|
||||
+ self.gamma * nextvalues * nextnonterminal
|
||||
- values_trajs[t]
|
||||
)
|
||||
# A = delta_t + gamma*lamdba*delta_{t+1} + ...
|
||||
advantages_trajs[t] = lastgaelam = (
|
||||
delta
|
||||
+ self.gamma
|
||||
* self.gae_lambda
|
||||
* nextnonterminal
|
||||
* lastgaelam
|
||||
)
|
||||
returns_trajs = advantages_trajs + values_trajs
|
||||
|
||||
# k for environment step
|
||||
obs_k = einops.rearrange(
|
||||
torch.tensor(obs_trajs).float().to(self.device),
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
chains_k = einops.rearrange(
|
||||
torch.tensor(chains_trajs).float().to(self.device),
|
||||
"s e t h d -> (s e) t h d",
|
||||
)
|
||||
returns_k = (
|
||||
torch.tensor(returns_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
values_k = (
|
||||
torch.tensor(values_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
advantages_k = (
|
||||
torch.tensor(advantages_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
logprobs_k = torch.tensor(logprobs_trajs).float().to(self.device)
|
||||
|
||||
# Update policy and critic
|
||||
total_steps = self.n_steps * self.n_envs
|
||||
inds_k = np.arange(total_steps)
|
||||
clipfracs = []
|
||||
for update_epoch in range(self.update_epochs):
|
||||
|
||||
# for each epoch, go through all data in batches
|
||||
flag_break = False
|
||||
np.random.shuffle(inds_k)
|
||||
num_batch = max(1, total_steps // self.batch_size) # skip last ones
|
||||
for batch in range(num_batch):
|
||||
start = batch * self.batch_size
|
||||
end = start + self.batch_size
|
||||
inds_b = inds_k[start:end] # b for batch
|
||||
obs_b = obs_k[inds_b]
|
||||
chains_b = chains_k[inds_b]
|
||||
returns_b = returns_k[inds_b]
|
||||
values_b = values_k[inds_b]
|
||||
advantages_b = advantages_k[inds_b]
|
||||
logprobs_b = logprobs_k[inds_b]
|
||||
|
||||
# get loss
|
||||
(
|
||||
pg_loss,
|
||||
entropy_loss,
|
||||
v_loss,
|
||||
clipfrac,
|
||||
approx_kl,
|
||||
ratio,
|
||||
bc_loss,
|
||||
eta,
|
||||
) = self.model.loss(
|
||||
obs_b,
|
||||
chains_b,
|
||||
returns_b,
|
||||
values_b,
|
||||
advantages_b,
|
||||
logprobs_b,
|
||||
use_bc_loss=self.use_bc_loss,
|
||||
reward_horizon=self.reward_horizon,
|
||||
)
|
||||
loss = (
|
||||
pg_loss
|
||||
+ entropy_loss * self.ent_coef
|
||||
+ v_loss * self.vf_coef
|
||||
+ bc_loss * self.bc_loss_coeff
|
||||
)
|
||||
clipfracs += [clipfrac]
|
||||
|
||||
# update policy and critic
|
||||
self.actor_optimizer.zero_grad()
|
||||
self.critic_optimizer.zero_grad()
|
||||
if self.learn_eta:
|
||||
self.eta_optimizer.zero_grad()
|
||||
loss.backward()
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
if self.max_grad_norm is not None:
|
||||
torch.nn.utils.clip_grad_norm_(
|
||||
self.model.actor_ft.parameters(), self.max_grad_norm
|
||||
)
|
||||
self.actor_optimizer.step()
|
||||
if self.learn_eta and batch % self.eta_update_interval == 0:
|
||||
self.eta_optimizer.step()
|
||||
self.critic_optimizer.step()
|
||||
log.info(
|
||||
f"approx_kl: {approx_kl}, update_epoch: {update_epoch}, num_batch: {num_batch}"
|
||||
)
|
||||
|
||||
# Stop gradient update if KL difference reaches target
|
||||
if self.target_kl is not None and approx_kl > self.target_kl:
|
||||
flag_break = True
|
||||
break
|
||||
if flag_break:
|
||||
break
|
||||
|
||||
# Explained variation of future rewards using value function
|
||||
y_pred, y_true = values_k.cpu().numpy(), returns_k.cpu().numpy()
|
||||
var_y = np.var(y_true)
|
||||
explained_var = (
|
||||
np.nan if var_y == 0 else 1 - np.var(y_true - y_pred) / var_y
|
||||
)
|
||||
|
||||
# Plot state trajectories in D3IL
|
||||
if (
|
||||
self.itr % self.render_freq == 0
|
||||
and self.n_render > 0
|
||||
and self.traj_plotter is not None
|
||||
):
|
||||
self.traj_plotter(
|
||||
obs_full_trajs=obs_full_trajs,
|
||||
n_render=self.n_render,
|
||||
max_episode_steps=self.max_episode_steps,
|
||||
render_dir=self.render_dir,
|
||||
itr=self.itr,
|
||||
)
|
||||
|
||||
# Update lr, min_sampling_std
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
self.actor_lr_scheduler.step()
|
||||
if self.learn_eta:
|
||||
self.eta_lr_scheduler.step()
|
||||
self.critic_lr_scheduler.step()
|
||||
self.model.step()
|
||||
diffusion_min_sampling_std = self.model.get_min_sampling_denoising_std()
|
||||
|
||||
# Save model
|
||||
if self.itr % self.save_model_freq == 0 or self.itr == self.n_train_itr - 1:
|
||||
self.save_model()
|
||||
|
||||
# Log loss and save metrics
|
||||
run_results.append(
|
||||
{
|
||||
"itr": self.itr,
|
||||
}
|
||||
)
|
||||
if self.save_trajs:
|
||||
run_results[-1]["obs_full_trajs"] = obs_full_trajs
|
||||
run_results[-1]["obs_trajs"] = obs_trajs
|
||||
run_results[-1]["chains_trajs"] = chains_trajs
|
||||
run_results[-1]["reward_trajs"] = reward_trajs
|
||||
if self.itr % self.log_freq == 0:
|
||||
time = timer()
|
||||
if eval_mode:
|
||||
log.info(
|
||||
f"eval: success rate {success_rate:8.4f} | avg episode reward {avg_episode_reward:8.4f} | avg best reward {avg_best_reward:8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"success rate - eval": success_rate,
|
||||
"avg episode reward - eval": avg_episode_reward,
|
||||
"avg best reward - eval": avg_best_reward,
|
||||
"num episode - eval": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=False,
|
||||
)
|
||||
run_results[-1]["eval_success_rate"] = success_rate
|
||||
run_results[-1]["eval_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["eval_best_reward"] = avg_best_reward
|
||||
else:
|
||||
log.info(
|
||||
f"{self.itr}: loss {loss:8.4f} | pg loss {pg_loss:8.4f} | value loss {v_loss:8.4f} | bc loss {bc_loss:8.4f} | reward {avg_episode_reward:8.4f} | eta {eta:8.4f} | t:{time:8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"loss": loss,
|
||||
"pg loss": pg_loss,
|
||||
"value loss": v_loss,
|
||||
"bc loss": bc_loss,
|
||||
"eta": eta,
|
||||
"approx kl": approx_kl,
|
||||
"ratio": ratio,
|
||||
"clipfrac": np.mean(clipfracs),
|
||||
"explained variance": explained_var,
|
||||
"avg episode reward - train": avg_episode_reward,
|
||||
"num episode - train": num_episode_finished,
|
||||
"diffusion - min sampling std": diffusion_min_sampling_std,
|
||||
"actor lr": self.actor_optimizer.param_groups[0]["lr"],
|
||||
"critic lr": self.critic_optimizer.param_groups[0][
|
||||
"lr"
|
||||
],
|
||||
},
|
||||
step=self.itr,
|
||||
commit=True,
|
||||
)
|
||||
run_results[-1]["loss"] = loss
|
||||
run_results[-1]["pg_loss"] = pg_loss
|
||||
run_results[-1]["value_loss"] = v_loss
|
||||
run_results[-1]["bc_loss"] = bc_loss
|
||||
run_results[-1]["eta"] = eta
|
||||
run_results[-1]["approx_kl"] = approx_kl
|
||||
run_results[-1]["ratio"] = ratio
|
||||
run_results[-1]["clip_frac"] = np.mean(clipfracs)
|
||||
run_results[-1]["explained_variance"] = explained_var
|
||||
run_results[-1]["train_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["time"] = time
|
||||
with open(self.result_path, "wb") as f:
|
||||
pickle.dump(run_results, f)
|
||||
self.itr += 1
|
||||
@@ -0,0 +1,468 @@
|
||||
"""
|
||||
DPPO fine-tuning for pixel observations.
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
import pickle
|
||||
import einops
|
||||
import numpy as np
|
||||
import torch
|
||||
import logging
|
||||
import wandb
|
||||
import math
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from util.timer import Timer
|
||||
from agent.finetune.train_ppo_diffusion_agent import TrainPPODiffusionAgent
|
||||
from model.common.modules import RandomShiftsAug
|
||||
|
||||
|
||||
class TrainPPOImgDiffusionAgent(TrainPPODiffusionAgent):
|
||||
|
||||
def __init__(self, cfg):
|
||||
super().__init__(cfg)
|
||||
|
||||
# Image randomization
|
||||
self.augment = cfg.train.augment
|
||||
if self.augment:
|
||||
self.aug = RandomShiftsAug(pad=4)
|
||||
|
||||
# Set obs dim - we will save the different obs in batch in a dict
|
||||
shape_meta = cfg.shape_meta
|
||||
self.obs_dims = {k: shape_meta.obs[k]["shape"] for k in shape_meta.obs.keys()}
|
||||
|
||||
# Gradient accumulation to deal with large GPU RAM usage
|
||||
self.grad_accumulate = cfg.train.grad_accumulate
|
||||
|
||||
def run(self):
|
||||
|
||||
# Start training loop
|
||||
timer = Timer()
|
||||
run_results = []
|
||||
last_itr_eval = False
|
||||
done_venv = np.zeros((1, self.n_envs))
|
||||
while self.itr < self.n_train_itr:
|
||||
|
||||
# Prepare video paths for each envs --- only applies for the first set of episodes if allowing reset within iteration and each iteration has multiple episodes from one env
|
||||
options_venv = [{} for _ in range(self.n_envs)]
|
||||
if self.itr % self.render_freq == 0 and self.render_video:
|
||||
for env_ind in range(self.n_render):
|
||||
options_venv[env_ind]["video_path"] = os.path.join(
|
||||
self.render_dir, f"itr-{self.itr}_trial-{env_ind}.mp4"
|
||||
)
|
||||
|
||||
# Define train or eval - all envs restart
|
||||
eval_mode = self.itr % self.val_freq == 0 and not self.force_train
|
||||
self.model.eval() if eval_mode else self.model.train()
|
||||
last_itr_eval = eval_mode
|
||||
|
||||
# Reset env before iteration starts (1) if specified, (2) at eval mode, or (3) right after eval mode
|
||||
dones_trajs = np.empty((0, self.n_envs))
|
||||
firsts_trajs = np.zeros((self.n_steps + 1, self.n_envs))
|
||||
if self.reset_at_iteration or eval_mode or last_itr_eval:
|
||||
prev_obs_venv = self.reset_env_all(options_venv=options_venv)
|
||||
firsts_trajs[0] = 1
|
||||
else:
|
||||
firsts_trajs[0] = (
|
||||
done_venv # if done at the end of last iteration, then the envs are just reset
|
||||
)
|
||||
|
||||
# Holder
|
||||
obs_trajs = {
|
||||
k: np.empty((0, self.n_envs, self.n_cond_step, *self.obs_dims[k]))
|
||||
for k in self.obs_dims
|
||||
}
|
||||
chains_trajs = np.empty(
|
||||
(
|
||||
0,
|
||||
self.n_envs,
|
||||
self.model.ft_denoising_steps + 1,
|
||||
self.horizon_steps,
|
||||
self.action_dim,
|
||||
)
|
||||
)
|
||||
reward_trajs = np.empty((0, self.n_envs))
|
||||
|
||||
# Collect a set of trajectories from env
|
||||
for step in range(self.n_steps):
|
||||
if step % 10 == 0:
|
||||
print(f"Processed step {step} of {self.n_steps}")
|
||||
|
||||
# Select action
|
||||
with torch.no_grad():
|
||||
cond = {
|
||||
key: torch.from_numpy(prev_obs_venv[key])
|
||||
.float()
|
||||
.to(self.device)
|
||||
for key in self.obs_dims.keys()
|
||||
} # batch each type of obs and put into dict
|
||||
samples = self.model(
|
||||
cond=cond,
|
||||
deterministic=eval_mode,
|
||||
return_chain=True,
|
||||
)
|
||||
output_venv = (
|
||||
samples.trajectories.cpu().numpy()
|
||||
) # n_env x horizon x act
|
||||
chains_venv = (
|
||||
samples.chains.cpu().numpy()
|
||||
) # n_env x denoising x horizon x act
|
||||
action_venv = output_venv[:, : self.act_steps]
|
||||
|
||||
# Apply multi-step action
|
||||
obs_venv, reward_venv, done_venv, info_venv = self.venv.step(
|
||||
action_venv
|
||||
)
|
||||
for k in obs_trajs.keys():
|
||||
obs_trajs[k] = np.vstack((obs_trajs[k], prev_obs_venv[k][None]))
|
||||
chains_trajs = np.vstack((chains_trajs, chains_venv[None]))
|
||||
reward_trajs = np.vstack((reward_trajs, reward_venv[None]))
|
||||
dones_trajs = np.vstack((dones_trajs, done_venv[None]))
|
||||
firsts_trajs[step + 1] = done_venv
|
||||
prev_obs_venv = obs_venv
|
||||
|
||||
# Summarize episode reward --- this needs to be handled differently depending on whether the environment is reset after each iteration. Only count episodes that finish within the iteration.
|
||||
episodes_start_end = []
|
||||
for env_ind in range(self.n_envs):
|
||||
env_steps = np.where(firsts_trajs[:, env_ind] == 1)[0]
|
||||
for i in range(len(env_steps) - 1):
|
||||
start = env_steps[i]
|
||||
end = env_steps[i + 1]
|
||||
if end - start > 1:
|
||||
episodes_start_end.append((env_ind, start, end - 1))
|
||||
if len(episodes_start_end) > 0:
|
||||
reward_trajs_split = [
|
||||
reward_trajs[start : end + 1, env_ind]
|
||||
for env_ind, start, end in episodes_start_end
|
||||
]
|
||||
num_episode_finished = len(reward_trajs_split)
|
||||
episode_reward = np.array(
|
||||
[np.sum(reward_traj) for reward_traj in reward_trajs_split]
|
||||
)
|
||||
episode_best_reward = np.array(
|
||||
[
|
||||
np.max(reward_traj) / self.act_steps
|
||||
for reward_traj in reward_trajs_split
|
||||
]
|
||||
)
|
||||
avg_episode_reward = np.mean(episode_reward)
|
||||
avg_best_reward = np.mean(episode_best_reward)
|
||||
success_rate = np.mean(
|
||||
episode_best_reward >= self.best_reward_threshold_for_success
|
||||
)
|
||||
else:
|
||||
episode_reward = np.array([])
|
||||
num_episode_finished = 0
|
||||
avg_episode_reward = 0
|
||||
avg_best_reward = 0
|
||||
success_rate = 0
|
||||
log.info("[WARNING] No episode completed within the iteration!")
|
||||
|
||||
# Update
|
||||
if not eval_mode:
|
||||
with torch.no_grad():
|
||||
# apply image randomization
|
||||
obs_trajs["rgb"] = (
|
||||
torch.from_numpy(obs_trajs["rgb"]).float().to(self.device)
|
||||
)
|
||||
obs_trajs["state"] = (
|
||||
torch.from_numpy(obs_trajs["state"]).float().to(self.device)
|
||||
)
|
||||
if self.augment:
|
||||
rgb = einops.rearrange(
|
||||
obs_trajs["rgb"],
|
||||
"s e t c h w -> (s e t) c h w",
|
||||
)
|
||||
rgb = self.aug(rgb)
|
||||
obs_trajs["rgb"] = einops.rearrange(
|
||||
rgb,
|
||||
"(s e t) c h w -> s e t c h w",
|
||||
s=self.n_steps,
|
||||
e=self.n_envs,
|
||||
)
|
||||
|
||||
# Calculate value and logprobs - split into batches to prevent out of memory
|
||||
num_split = math.ceil(
|
||||
self.n_envs * self.n_steps / self.logprob_batch_size
|
||||
)
|
||||
obs_ts = [{} for _ in range(num_split)]
|
||||
for k in obs_trajs.keys():
|
||||
obs_k = einops.rearrange(
|
||||
obs_trajs[k],
|
||||
"s e ... -> (s e) ...",
|
||||
)
|
||||
obs_ts_k = torch.split(obs_k, self.logprob_batch_size, dim=0)
|
||||
for i, obs_t in enumerate(obs_ts_k):
|
||||
obs_ts[i][k] = obs_t
|
||||
values_trajs = np.empty((0, self.n_envs))
|
||||
for obs in obs_ts:
|
||||
values = (
|
||||
self.model.critic(obs, no_augment=True)
|
||||
.cpu()
|
||||
.numpy()
|
||||
.flatten()
|
||||
)
|
||||
values_trajs = np.vstack(
|
||||
(values_trajs, values.reshape(-1, self.n_envs))
|
||||
)
|
||||
chains_t = einops.rearrange(
|
||||
torch.from_numpy(chains_trajs).float().to(self.device),
|
||||
"s e t h d -> (s e) t h d",
|
||||
)
|
||||
chains_ts = torch.split(chains_t, self.logprob_batch_size, dim=0)
|
||||
logprobs_trajs = np.empty(
|
||||
(
|
||||
0,
|
||||
self.model.ft_denoising_steps,
|
||||
self.horizon_steps,
|
||||
self.action_dim,
|
||||
)
|
||||
)
|
||||
for obs, chains in zip(obs_ts, chains_ts):
|
||||
logprobs = self.model.get_logprobs(obs, chains).cpu().numpy()
|
||||
logprobs_trajs = np.vstack(
|
||||
(
|
||||
logprobs_trajs,
|
||||
logprobs.reshape(-1, *logprobs_trajs.shape[1:]),
|
||||
)
|
||||
)
|
||||
|
||||
# normalize reward with running variance if specified
|
||||
if self.reward_scale_running:
|
||||
reward_trajs_transpose = self.running_reward_scaler(
|
||||
reward=reward_trajs.T, first=firsts_trajs[:-1].T
|
||||
)
|
||||
reward_trajs = reward_trajs_transpose.T
|
||||
|
||||
# bootstrap value with GAE if not done - apply reward scaling with constant if specified
|
||||
obs_venv_ts = {
|
||||
key: torch.from_numpy(obs_venv[key]).float().to(self.device)
|
||||
for key in self.obs_dims.keys()
|
||||
}
|
||||
with torch.no_grad():
|
||||
next_value = (
|
||||
self.model.critic(obs_venv_ts, no_augment=True)
|
||||
.reshape(1, -1)
|
||||
.cpu()
|
||||
.numpy()
|
||||
)
|
||||
advantages_trajs = np.zeros_like(reward_trajs)
|
||||
lastgaelam = 0
|
||||
for t in reversed(range(self.n_steps)):
|
||||
if t == self.n_steps - 1:
|
||||
nextnonterminal = 1.0 - done_venv
|
||||
nextvalues = next_value
|
||||
else:
|
||||
nextnonterminal = 1.0 - dones_trajs[t + 1]
|
||||
nextvalues = values_trajs[t + 1]
|
||||
# delta = r + gamma*V(st+1) - V(st)
|
||||
delta = (
|
||||
reward_trajs[t] * self.reward_scale_const
|
||||
+ self.gamma * nextvalues * nextnonterminal
|
||||
- values_trajs[t]
|
||||
)
|
||||
# A = delta_t + gamma*lamdba*delta_{t+1} + ...
|
||||
advantages_trajs[t] = lastgaelam = (
|
||||
delta
|
||||
+ self.gamma
|
||||
* self.gae_lambda
|
||||
* nextnonterminal
|
||||
* lastgaelam
|
||||
)
|
||||
returns_trajs = advantages_trajs + values_trajs
|
||||
|
||||
# k for environment step
|
||||
obs_k = {
|
||||
k: einops.rearrange(
|
||||
obs_trajs[k],
|
||||
"s e ... -> (s e) ...",
|
||||
)
|
||||
for k in obs_trajs.keys()
|
||||
}
|
||||
chains_k = einops.rearrange(
|
||||
torch.tensor(chains_trajs).float().to(self.device),
|
||||
"s e t h d -> (s e) t h d",
|
||||
)
|
||||
returns_k = (
|
||||
torch.tensor(returns_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
values_k = (
|
||||
torch.tensor(values_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
advantages_k = (
|
||||
torch.tensor(advantages_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
logprobs_k = torch.tensor(logprobs_trajs).float().to(self.device)
|
||||
|
||||
# Update policy and critic
|
||||
total_steps = self.n_steps * self.n_envs
|
||||
inds_k = np.arange(total_steps)
|
||||
clipfracs = []
|
||||
for update_epoch in range(self.update_epochs):
|
||||
|
||||
# for each epoch, go through all data in batches
|
||||
flag_break = False
|
||||
np.random.shuffle(inds_k)
|
||||
num_batch = max(1, total_steps // self.batch_size) # skip last ones
|
||||
for batch in range(num_batch):
|
||||
start = batch * self.batch_size
|
||||
end = start + self.batch_size
|
||||
inds_b = inds_k[start:end] # b for batch
|
||||
obs_b = {k: obs_k[k][inds_b] for k in obs_k.keys()}
|
||||
chains_b = chains_k[inds_b]
|
||||
returns_b = returns_k[inds_b]
|
||||
values_b = values_k[inds_b]
|
||||
advantages_b = advantages_k[inds_b]
|
||||
logprobs_b = logprobs_k[inds_b]
|
||||
|
||||
# get loss
|
||||
(
|
||||
pg_loss,
|
||||
entropy_loss,
|
||||
v_loss,
|
||||
clipfrac,
|
||||
approx_kl,
|
||||
ratio,
|
||||
bc_loss,
|
||||
eta,
|
||||
) = self.model.loss(
|
||||
obs_b,
|
||||
chains_b,
|
||||
returns_b,
|
||||
values_b,
|
||||
advantages_b,
|
||||
logprobs_b,
|
||||
use_bc_loss=self.use_bc_loss,
|
||||
reward_horizon=self.reward_horizon,
|
||||
)
|
||||
loss = (
|
||||
pg_loss
|
||||
+ entropy_loss * self.ent_coef
|
||||
+ v_loss * self.vf_coef
|
||||
+ bc_loss * self.bc_loss_coeff
|
||||
)
|
||||
clipfracs += [clipfrac]
|
||||
|
||||
# update policy and critic
|
||||
loss.backward()
|
||||
if (batch + 1) % self.grad_accumulate == 0:
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
if self.max_grad_norm is not None:
|
||||
torch.nn.utils.clip_grad_norm_(
|
||||
self.model.actor_ft.parameters(),
|
||||
self.max_grad_norm,
|
||||
)
|
||||
self.actor_optimizer.step()
|
||||
if (
|
||||
self.learn_eta
|
||||
and batch % self.eta_update_interval == 0
|
||||
):
|
||||
self.eta_optimizer.step()
|
||||
self.critic_optimizer.step()
|
||||
self.actor_optimizer.zero_grad()
|
||||
self.critic_optimizer.zero_grad()
|
||||
if self.learn_eta:
|
||||
self.eta_optimizer.zero_grad()
|
||||
log.info(f"run grad update at batch {batch}")
|
||||
log.info(
|
||||
f"approx_kl: {approx_kl}, update_epoch: {update_epoch}, num_batch: {num_batch}"
|
||||
)
|
||||
|
||||
# Stop gradient update if KL difference reaches target
|
||||
if (
|
||||
self.target_kl is not None
|
||||
and approx_kl > self.target_kl
|
||||
and self.itr >= self.n_critic_warmup_itr
|
||||
):
|
||||
flag_break = True
|
||||
break
|
||||
if flag_break:
|
||||
break
|
||||
|
||||
# Explained variation of future rewards using value function
|
||||
y_pred, y_true = values_k.cpu().numpy(), returns_k.cpu().numpy()
|
||||
var_y = np.var(y_true)
|
||||
explained_var = (
|
||||
np.nan if var_y == 0 else 1 - np.var(y_true - y_pred) / var_y
|
||||
)
|
||||
|
||||
# Update lr
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
self.actor_lr_scheduler.step()
|
||||
if self.learn_eta:
|
||||
self.eta_lr_scheduler.step()
|
||||
self.critic_lr_scheduler.step()
|
||||
self.model.step()
|
||||
diffusion_min_sampling_std = self.model.get_min_sampling_denoising_std()
|
||||
|
||||
# Save model
|
||||
if self.itr % self.save_model_freq == 0 or self.itr == self.n_train_itr - 1:
|
||||
self.save_model()
|
||||
|
||||
# Log loss and save metrics
|
||||
run_results.append(
|
||||
{
|
||||
"itr": self.itr,
|
||||
}
|
||||
)
|
||||
if self.itr % self.log_freq == 0:
|
||||
if eval_mode:
|
||||
log.info(
|
||||
f"eval: success rate {success_rate:8.4f} | avg episode reward {avg_episode_reward:8.4f} | avg best reward {avg_best_reward:8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"success rate - eval": success_rate,
|
||||
"avg episode reward - eval": avg_episode_reward,
|
||||
"avg best reward - eval": avg_best_reward,
|
||||
"num episode - eval": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=False,
|
||||
)
|
||||
run_results[-1]["eval_success_rate"] = success_rate
|
||||
run_results[-1]["eval_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["eval_best_reward"] = avg_best_reward
|
||||
else:
|
||||
log.info(
|
||||
f"{self.itr}: loss {loss:8.4f} | pg loss {pg_loss:8.4f} | value loss {v_loss:8.4f} | bc loss {bc_loss:8.4f} | reward {avg_episode_reward:8.4f} | eta {eta:8.4f} | t:{timer():8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"loss": loss,
|
||||
"pg loss": pg_loss,
|
||||
"value loss": v_loss,
|
||||
"bc loss": bc_loss,
|
||||
"eta": eta,
|
||||
"approx kl": approx_kl,
|
||||
"ratio": ratio,
|
||||
"clipfrac": np.mean(clipfracs),
|
||||
"explained variance": explained_var,
|
||||
"avg episode reward - train": avg_episode_reward,
|
||||
"num episode - train": num_episode_finished,
|
||||
"diffusion - min sampling std": diffusion_min_sampling_std,
|
||||
"actor lr": self.actor_optimizer.param_groups[0]["lr"],
|
||||
"critic lr": self.critic_optimizer.param_groups[0][
|
||||
"lr"
|
||||
],
|
||||
},
|
||||
step=self.itr,
|
||||
commit=True,
|
||||
)
|
||||
run_results[-1]["loss"] = loss
|
||||
run_results[-1]["pg_loss"] = pg_loss
|
||||
run_results[-1]["value_loss"] = v_loss
|
||||
run_results[-1]["bc_loss"] = bc_loss
|
||||
run_results[-1]["eta"] = eta
|
||||
run_results[-1]["approx_kl"] = approx_kl
|
||||
run_results[-1]["ratio"] = ratio
|
||||
run_results[-1]["clip_frac"] = np.mean(clipfracs)
|
||||
run_results[-1]["explained_variance"] = explained_var
|
||||
run_results[-1]["train_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["time"] = timer()
|
||||
with open(self.result_path, "wb") as f:
|
||||
pickle.dump(run_results, f)
|
||||
self.itr += 1
|
||||
@@ -0,0 +1,405 @@
|
||||
"""
|
||||
Use diffusion exact likelihood for policy gradient.
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
import pickle
|
||||
import einops
|
||||
import numpy as np
|
||||
import torch
|
||||
import logging
|
||||
import wandb
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from util.timer import Timer
|
||||
from agent.finetune.train_ppo_diffusion_agent import TrainPPODiffusionAgent
|
||||
|
||||
|
||||
class TrainPPOExactDiffusionAgent(TrainPPODiffusionAgent):
|
||||
|
||||
def __init__(self, cfg):
|
||||
super().__init__(cfg)
|
||||
|
||||
def run(self):
|
||||
"""
|
||||
For exact likelihood, we do not need to save the chains.
|
||||
"""
|
||||
|
||||
# Start training loop
|
||||
timer = Timer()
|
||||
run_results = []
|
||||
last_itr_eval = False
|
||||
done_venv = np.zeros((1, self.n_envs))
|
||||
while self.itr < self.n_train_itr:
|
||||
|
||||
# Prepare video paths for each envs --- only applies for the first set of episodes if allowing reset within iteration and each iteration has multiple episodes from one env
|
||||
options_venv = [{} for _ in range(self.n_envs)]
|
||||
if self.itr % self.render_freq == 0 and self.render_video:
|
||||
for env_ind in range(self.n_render):
|
||||
options_venv[env_ind]["video_path"] = os.path.join(
|
||||
self.render_dir, f"itr-{self.itr}_trial-{env_ind}.mp4"
|
||||
)
|
||||
|
||||
# Define train or eval - all envs restart
|
||||
eval_mode = self.itr % self.val_freq == 0 and not self.force_train
|
||||
self.model.eval() if eval_mode else self.model.train()
|
||||
last_itr_eval = eval_mode
|
||||
|
||||
# Reset env before iteration starts (1) if specified, (2) at eval mode, or (3) right after eval mode
|
||||
dones_trajs = np.empty((0, self.n_envs))
|
||||
firsts_trajs = np.zeros((self.n_steps + 1, self.n_envs))
|
||||
if self.reset_at_iteration or eval_mode or last_itr_eval:
|
||||
prev_obs_venv = self.reset_env_all(options_venv=options_venv)
|
||||
firsts_trajs[0] = 1
|
||||
else:
|
||||
firsts_trajs[0] = (
|
||||
done_venv # if done at the end of last iteration, then the envs are just reset
|
||||
)
|
||||
|
||||
# Holder
|
||||
obs_trajs = np.empty((0, self.n_envs, self.n_cond_step, self.obs_dim))
|
||||
samples_trajs = np.empty(
|
||||
(
|
||||
0,
|
||||
self.n_envs,
|
||||
self.horizon_steps,
|
||||
self.action_dim,
|
||||
)
|
||||
)
|
||||
chains_trajs = np.empty(
|
||||
(
|
||||
0,
|
||||
self.n_envs,
|
||||
self.model.ft_denoising_steps + 1,
|
||||
self.horizon_steps,
|
||||
self.action_dim,
|
||||
)
|
||||
)
|
||||
reward_trajs = np.empty((0, self.n_envs))
|
||||
obs_full_trajs = np.empty((0, self.n_envs, self.obs_dim))
|
||||
obs_full_trajs = np.vstack(
|
||||
(obs_full_trajs, prev_obs_venv[None].squeeze(2))
|
||||
) # remove cond_step dim
|
||||
|
||||
# Collect a set of trajectories from env
|
||||
for step in range(self.n_steps):
|
||||
if step % 10 == 0:
|
||||
print(f"Processed step {step} of {self.n_steps}")
|
||||
|
||||
# Select action
|
||||
with torch.no_grad():
|
||||
samples = self.model(
|
||||
cond=torch.from_numpy(prev_obs_venv).float().to(self.device),
|
||||
deterministic=eval_mode,
|
||||
return_chain=True,
|
||||
)
|
||||
output_venv = (
|
||||
samples.trajectories.cpu().numpy()
|
||||
) # n_env x horizon x act
|
||||
chains_venv = (
|
||||
samples.chains.cpu().numpy()
|
||||
) # n_env x denoising x horizon x act
|
||||
action_venv = output_venv[:, : self.act_steps]
|
||||
|
||||
# Apply multi-step action
|
||||
obs_venv, reward_venv, done_venv, info_venv = self.venv.step(
|
||||
action_venv
|
||||
)
|
||||
if self.save_full_observations:
|
||||
obs_full_venv = np.vstack(
|
||||
[info["full_obs"][None] for info in info_venv]
|
||||
) # n_envs x n_act_steps x obs_dim
|
||||
obs_full_trajs = np.vstack(
|
||||
(obs_full_trajs, obs_full_venv.transpose(1, 0, 2))
|
||||
)
|
||||
obs_trajs = np.vstack((obs_trajs, prev_obs_venv[None]))
|
||||
chains_trajs = np.vstack((chains_trajs, chains_venv[None]))
|
||||
samples_trajs = np.vstack((samples_trajs, output_venv[None]))
|
||||
reward_trajs = np.vstack((reward_trajs, reward_venv[None]))
|
||||
dones_trajs = np.vstack((dones_trajs, done_venv[None]))
|
||||
firsts_trajs[step + 1] = done_venv
|
||||
prev_obs_venv = obs_venv
|
||||
|
||||
# Summarize episode reward --- this needs to be handled differently depending on whether the environment is reset after each iteration. Only count episodes that finish within the iteration.
|
||||
episodes_start_end = []
|
||||
for env_ind in range(self.n_envs):
|
||||
env_steps = np.where(firsts_trajs[:, env_ind] == 1)[0]
|
||||
for i in range(len(env_steps) - 1):
|
||||
start = env_steps[i]
|
||||
end = env_steps[i + 1]
|
||||
if end - start > 1:
|
||||
episodes_start_end.append((env_ind, start, end - 1))
|
||||
if len(episodes_start_end) > 0:
|
||||
reward_trajs_split = [
|
||||
reward_trajs[start : end + 1, env_ind]
|
||||
for env_ind, start, end in episodes_start_end
|
||||
]
|
||||
num_episode_finished = len(reward_trajs_split)
|
||||
episode_reward = np.array(
|
||||
[np.sum(reward_traj) for reward_traj in reward_trajs_split]
|
||||
)
|
||||
episode_best_reward = np.array(
|
||||
[
|
||||
np.max(reward_traj) / self.act_steps
|
||||
for reward_traj in reward_trajs_split
|
||||
]
|
||||
)
|
||||
avg_episode_reward = np.mean(episode_reward)
|
||||
avg_best_reward = np.mean(episode_best_reward)
|
||||
success_rate = np.mean(
|
||||
episode_best_reward >= self.best_reward_threshold_for_success
|
||||
)
|
||||
else:
|
||||
episode_reward = np.array([])
|
||||
num_episode_finished = 0
|
||||
avg_episode_reward = 0
|
||||
avg_best_reward = 0
|
||||
success_rate = 0
|
||||
log.info("[WARNING] No episode completed within the iteration!")
|
||||
|
||||
# Update
|
||||
if not eval_mode:
|
||||
with torch.no_grad():
|
||||
# Calculate value and logprobs - split into batches to prevent out of memory
|
||||
obs_t = einops.rearrange(
|
||||
torch.from_numpy(obs_trajs).float().to(self.device),
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
obs_ts = torch.split(obs_t, self.logprob_batch_size, dim=0)
|
||||
values_trajs = np.empty((0, self.n_envs))
|
||||
for obs in obs_ts:
|
||||
values = self.model.critic(obs).cpu().numpy().flatten()
|
||||
values_trajs = np.vstack(
|
||||
(values_trajs, values.reshape(-1, self.n_envs))
|
||||
)
|
||||
samples_t = einops.rearrange(
|
||||
torch.from_numpy(samples_trajs).float().to(self.device),
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
samples_ts = torch.split(samples_t, self.logprob_batch_size, dim=0)
|
||||
logprobs_trajs = np.empty((0))
|
||||
for obs, samples in zip(obs_ts, samples_ts):
|
||||
logprobs = (
|
||||
self.model.get_exact_logprobs(obs, samples).cpu().numpy()
|
||||
)
|
||||
logprobs_trajs = np.concatenate((logprobs_trajs, logprobs))
|
||||
|
||||
# normalize reward with running variance if specified
|
||||
if self.reward_scale_running:
|
||||
reward_trajs_transpose = self.running_reward_scaler(
|
||||
reward=reward_trajs.T, first=firsts_trajs[:-1].T
|
||||
)
|
||||
reward_trajs = reward_trajs_transpose.T
|
||||
|
||||
# bootstrap value with GAE if not done - apply reward scaling with constant if specified
|
||||
obs_venv_ts = torch.from_numpy(obs_venv).float().to(self.device)
|
||||
with torch.no_grad():
|
||||
next_value = (
|
||||
self.model.critic(obs_venv_ts).reshape(1, -1).cpu().numpy()
|
||||
)
|
||||
advantages_trajs = np.zeros_like(reward_trajs)
|
||||
lastgaelam = 0
|
||||
for t in reversed(range(self.n_steps)):
|
||||
if t == self.n_steps - 1:
|
||||
nextnonterminal = 1.0 - done_venv
|
||||
nextvalues = next_value
|
||||
else:
|
||||
nextnonterminal = 1.0 - dones_trajs[t + 1]
|
||||
nextvalues = values_trajs[t + 1]
|
||||
# delta = r + gamma*V(st+1) - V(st)
|
||||
delta = (
|
||||
reward_trajs[t] * self.reward_scale_const
|
||||
+ self.gamma * nextvalues * nextnonterminal
|
||||
- values_trajs[t]
|
||||
)
|
||||
# A = delta_t + gamma*lamdba*delta_{t+1} + ...
|
||||
advantages_trajs[t] = lastgaelam = (
|
||||
delta
|
||||
+ self.gamma
|
||||
* self.gae_lambda
|
||||
* nextnonterminal
|
||||
* lastgaelam
|
||||
)
|
||||
returns_trajs = advantages_trajs + values_trajs
|
||||
|
||||
# k for environment step
|
||||
obs_k = einops.rearrange(
|
||||
torch.tensor(obs_trajs).float().to(self.device),
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
samples_k = einops.rearrange(
|
||||
torch.tensor(samples_trajs).float().to(self.device),
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
returns_k = (
|
||||
torch.tensor(returns_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
values_k = (
|
||||
torch.tensor(values_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
advantages_k = (
|
||||
torch.tensor(advantages_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
logprobs_k = torch.tensor(logprobs_trajs).float().to(self.device)
|
||||
|
||||
# Update policy and critic
|
||||
total_steps = self.n_steps * self.n_envs
|
||||
inds_k = np.arange(total_steps)
|
||||
clipfracs = []
|
||||
for update_epoch in range(self.update_epochs):
|
||||
|
||||
# for each epoch, go through all data in batches
|
||||
flag_break = False
|
||||
np.random.shuffle(inds_k)
|
||||
num_batch = max(1, total_steps // self.batch_size) # skip last ones
|
||||
for batch in range(num_batch):
|
||||
start = batch * self.batch_size
|
||||
end = start + self.batch_size
|
||||
inds_b = inds_k[start:end] # b for batch
|
||||
obs_b = obs_k[inds_b]
|
||||
samples_b = samples_k[inds_b]
|
||||
returns_b = returns_k[inds_b]
|
||||
values_b = values_k[inds_b]
|
||||
advantages_b = advantages_k[inds_b]
|
||||
logprobs_b = logprobs_k[inds_b]
|
||||
|
||||
# get loss
|
||||
(
|
||||
pg_loss,
|
||||
v_loss,
|
||||
clipfrac,
|
||||
approx_kl,
|
||||
ratio,
|
||||
bc_loss,
|
||||
) = self.model.loss(
|
||||
obs_b,
|
||||
samples_b,
|
||||
returns_b,
|
||||
values_b,
|
||||
advantages_b,
|
||||
logprobs_b,
|
||||
use_bc_loss=self.use_bc_loss,
|
||||
reward_horizon=self.reward_horizon,
|
||||
)
|
||||
loss = (
|
||||
pg_loss
|
||||
+ v_loss * self.vf_coef
|
||||
+ bc_loss * self.bc_loss_coeff
|
||||
)
|
||||
clipfracs += [clipfrac]
|
||||
|
||||
# update policy and critic
|
||||
self.actor_optimizer.zero_grad()
|
||||
self.critic_optimizer.zero_grad()
|
||||
loss.backward()
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
if self.max_grad_norm is not None:
|
||||
torch.nn.utils.clip_grad_norm_(
|
||||
self.model.actor_ft.parameters(), self.max_grad_norm
|
||||
)
|
||||
self.actor_optimizer.step()
|
||||
self.critic_optimizer.step()
|
||||
log.info(
|
||||
f"approx_kl: {approx_kl}, update_epoch: {update_epoch}, num_batch: {num_batch}"
|
||||
)
|
||||
|
||||
# Stop gradient update if KL difference reaches target
|
||||
if self.target_kl is not None and approx_kl > self.target_kl:
|
||||
flag_break = True
|
||||
break
|
||||
if flag_break:
|
||||
break
|
||||
|
||||
# Explained variation of future rewards using value function
|
||||
y_pred, y_true = values_k.cpu().numpy(), returns_k.cpu().numpy()
|
||||
var_y = np.var(y_true)
|
||||
explained_var = (
|
||||
np.nan if var_y == 0 else 1 - np.var(y_true - y_pred) / var_y
|
||||
)
|
||||
|
||||
# Plot state trajectories
|
||||
if (
|
||||
self.itr % self.render_freq == 0
|
||||
and self.n_render > 0
|
||||
and self.traj_plotter is not None
|
||||
):
|
||||
self.traj_plotter(
|
||||
obs_full_trajs=obs_full_trajs,
|
||||
n_render=self.n_render,
|
||||
max_episode_steps=self.max_episode_steps,
|
||||
render_dir=self.render_dir,
|
||||
itr=self.itr,
|
||||
)
|
||||
|
||||
# Update lr
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
self.actor_lr_scheduler.step()
|
||||
self.critic_lr_scheduler.step()
|
||||
|
||||
# Save model
|
||||
if self.itr % self.save_model_freq == 0 or self.itr == self.n_train_itr - 1:
|
||||
self.save_model()
|
||||
|
||||
# Log loss and save metrics
|
||||
run_results.append(
|
||||
{
|
||||
"itr": self.itr,
|
||||
}
|
||||
)
|
||||
if self.save_trajs:
|
||||
run_results[-1]["obs_full_trajs"] = obs_full_trajs
|
||||
run_results[-1]["obs_trajs"] = obs_trajs
|
||||
run_results[-1]["action_trajs"] = samples_trajs
|
||||
run_results[-1]["chains_trajs"] = chains_trajs
|
||||
run_results[-1]["reward_trajs"] = reward_trajs
|
||||
if self.itr % self.log_freq == 0:
|
||||
if eval_mode:
|
||||
log.info(
|
||||
f"eval: success rate {success_rate:8.4f} | avg episode reward {avg_episode_reward:8.4f} | avg best reward {avg_best_reward:8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"success rate - eval": success_rate,
|
||||
"avg episode reward - eval": avg_episode_reward,
|
||||
"avg best reward - eval": avg_best_reward,
|
||||
"num episode - eval": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=False,
|
||||
)
|
||||
run_results[-1]["eval_success_rate"] = success_rate
|
||||
run_results[-1]["eval_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["eval_best_reward"] = avg_best_reward
|
||||
else:
|
||||
log.info(
|
||||
f"{self.itr}: loss {loss:8.4f} | pg loss {pg_loss:8.4f} | value loss {v_loss:8.4f} | reward {avg_episode_reward:8.4f} | t:{timer():8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"loss": loss,
|
||||
"pg loss": pg_loss,
|
||||
"value loss": v_loss,
|
||||
"approx kl": approx_kl,
|
||||
"ratio": ratio,
|
||||
"clipfrac": np.mean(clipfracs),
|
||||
"explained variance": explained_var,
|
||||
"avg episode reward - train": avg_episode_reward,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=True,
|
||||
)
|
||||
run_results[-1]["loss"] = loss
|
||||
run_results[-1]["pg_loss"] = pg_loss
|
||||
run_results[-1]["value_loss"] = v_loss
|
||||
run_results[-1]["approx_kl"] = approx_kl
|
||||
run_results[-1]["ratio"] = ratio
|
||||
run_results[-1]["clip_frac"] = np.mean(clipfracs)
|
||||
run_results[-1]["explained_variance"] = explained_var
|
||||
run_results[-1]["train_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["time"] = timer()
|
||||
with open(self.result_path, "wb") as f:
|
||||
pickle.dump(run_results, f)
|
||||
self.itr += 1
|
||||
@@ -0,0 +1,404 @@
|
||||
"""
|
||||
PPO training for Gaussian/GMM policy.
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
import pickle
|
||||
import einops
|
||||
import numpy as np
|
||||
import torch
|
||||
import logging
|
||||
import wandb
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from util.timer import Timer
|
||||
from agent.finetune.train_ppo_agent import TrainPPOAgent
|
||||
|
||||
|
||||
class TrainPPOGaussianAgent(TrainPPOAgent):
|
||||
|
||||
def __init__(self, cfg):
|
||||
super().__init__(cfg)
|
||||
|
||||
def run(self):
|
||||
|
||||
# Start training loop
|
||||
timer = Timer()
|
||||
run_results = []
|
||||
last_itr_eval = False
|
||||
done_venv = np.zeros((1, self.n_envs))
|
||||
while self.itr < self.n_train_itr:
|
||||
|
||||
# Prepare video paths for each envs --- only applies for the first set of episodes if allowing reset within iteration and each iteration has multiple episodes from one env
|
||||
options_venv = [{} for _ in range(self.n_envs)]
|
||||
if self.itr % self.render_freq == 0 and self.render_video:
|
||||
for env_ind in range(self.n_render):
|
||||
options_venv[env_ind]["video_path"] = os.path.join(
|
||||
self.render_dir, f"itr-{self.itr}_trial-{env_ind}.mp4"
|
||||
)
|
||||
|
||||
# Define train or eval - all envs restart
|
||||
eval_mode = self.itr % self.val_freq == 0 and not self.force_train
|
||||
self.model.eval() if eval_mode else self.model.train()
|
||||
last_itr_eval = eval_mode
|
||||
|
||||
# Reset env before iteration starts (1) if specified, (2) at eval mode, or (3) right after eval mode
|
||||
dones_trajs = np.empty((0, self.n_envs))
|
||||
firsts_trajs = np.zeros((self.n_steps + 1, self.n_envs))
|
||||
if self.reset_at_iteration or eval_mode or last_itr_eval:
|
||||
prev_obs_venv = self.reset_env_all(options_venv=options_venv)
|
||||
firsts_trajs[0] = 1
|
||||
else:
|
||||
firsts_trajs[0] = (
|
||||
done_venv # if done at the end of last iteration, then the envs are just reset
|
||||
)
|
||||
|
||||
# Holder
|
||||
obs_trajs = np.empty((0, self.n_envs, self.n_cond_step, self.obs_dim))
|
||||
samples_trajs = np.empty(
|
||||
(
|
||||
0,
|
||||
self.n_envs,
|
||||
self.horizon_steps,
|
||||
self.action_dim,
|
||||
)
|
||||
)
|
||||
reward_trajs = np.empty((0, self.n_envs))
|
||||
obs_full_trajs = np.empty((0, self.n_envs, self.obs_dim))
|
||||
obs_full_trajs = np.vstack(
|
||||
(obs_full_trajs, prev_obs_venv[None].squeeze(2))
|
||||
) # remove cond_step dim
|
||||
|
||||
# Collect a set of trajectories from env
|
||||
for step in range(self.n_steps):
|
||||
if step % 10 == 0:
|
||||
print(f"Processed step {step} of {self.n_steps}")
|
||||
|
||||
# Select action
|
||||
with torch.no_grad():
|
||||
samples = self.model(
|
||||
cond=torch.from_numpy(prev_obs_venv).float().to(self.device),
|
||||
deterministic=eval_mode,
|
||||
)
|
||||
output_venv = samples.cpu().numpy()
|
||||
action_venv = output_venv[:, : self.act_steps, : self.action_dim]
|
||||
obs_trajs = np.vstack((obs_trajs, prev_obs_venv[None]))
|
||||
samples_trajs = np.vstack((samples_trajs, output_venv[None]))
|
||||
|
||||
# Apply multi-step action
|
||||
obs_venv, reward_venv, done_venv, info_venv = self.venv.step(
|
||||
action_venv
|
||||
)
|
||||
if self.save_full_observations:
|
||||
obs_full_venv = np.vstack(
|
||||
[info["full_obs"][None] for info in info_venv]
|
||||
) # n_envs x n_act_steps x obs_dim
|
||||
obs_full_trajs = np.vstack(
|
||||
(obs_full_trajs, obs_full_venv.transpose(1, 0, 2))
|
||||
)
|
||||
reward_trajs = np.vstack((reward_trajs, reward_venv[None]))
|
||||
dones_trajs = np.vstack((dones_trajs, done_venv[None]))
|
||||
firsts_trajs[step + 1] = done_venv
|
||||
prev_obs_venv = obs_venv
|
||||
|
||||
# Summarize episode reward --- this needs to be handled differently depending on whether the environment is reset after each iteration. Only count episodes that finish within the iteration.
|
||||
episodes_start_end = []
|
||||
for env_ind in range(self.n_envs):
|
||||
env_steps = np.where(firsts_trajs[:, env_ind] == 1)[0]
|
||||
for i in range(len(env_steps) - 1):
|
||||
start = env_steps[i]
|
||||
end = env_steps[i + 1]
|
||||
if end - start > 1:
|
||||
episodes_start_end.append((env_ind, start, end - 1))
|
||||
if len(episodes_start_end) > 0:
|
||||
reward_trajs_split = [
|
||||
reward_trajs[start : end + 1, env_ind]
|
||||
for env_ind, start, end in episodes_start_end
|
||||
]
|
||||
num_episode_finished = len(reward_trajs_split)
|
||||
episode_reward = np.array(
|
||||
[np.sum(reward_traj) for reward_traj in reward_trajs_split]
|
||||
)
|
||||
if (
|
||||
self.furniture_sparse_reward
|
||||
): # only for furniture tasks, where reward only occurs in one env step
|
||||
episode_best_reward = episode_reward
|
||||
else:
|
||||
episode_best_reward = np.array(
|
||||
[
|
||||
np.max(reward_traj) / self.act_steps
|
||||
for reward_traj in reward_trajs_split
|
||||
]
|
||||
)
|
||||
avg_episode_reward = np.mean(episode_reward)
|
||||
avg_best_reward = np.mean(episode_best_reward)
|
||||
success_rate = np.mean(
|
||||
episode_best_reward >= self.best_reward_threshold_for_success
|
||||
)
|
||||
else:
|
||||
episode_reward = np.array([])
|
||||
num_episode_finished = 0
|
||||
avg_episode_reward = 0
|
||||
avg_best_reward = 0
|
||||
success_rate = 0
|
||||
log.info("[WARNING] No episode completed within the iteration!")
|
||||
|
||||
# Update
|
||||
if not eval_mode:
|
||||
with torch.no_grad():
|
||||
# Calculate value and logprobs - split into batches to prevent out of memory
|
||||
obs_t = einops.rearrange(
|
||||
torch.from_numpy(obs_trajs).float().to(self.device),
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
obs_ts = torch.split(obs_t, self.logprob_batch_size, dim=0)
|
||||
values_trajs = np.empty((0, self.n_envs))
|
||||
for obs in obs_ts:
|
||||
values = self.model.critic(obs).cpu().numpy().flatten()
|
||||
values_trajs = np.vstack(
|
||||
(values_trajs, values.reshape(-1, self.n_envs))
|
||||
)
|
||||
samples_t = einops.rearrange(
|
||||
torch.from_numpy(samples_trajs).float().to(self.device),
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
samples_ts = torch.split(samples_t, self.logprob_batch_size, dim=0)
|
||||
logprobs_trajs = np.empty((0))
|
||||
for obs_t, samples_t in zip(obs_ts, samples_ts):
|
||||
logprobs = (
|
||||
self.model.get_logprobs(obs_t, samples_t)[0].cpu().numpy()
|
||||
)
|
||||
logprobs_trajs = np.concatenate(
|
||||
(
|
||||
logprobs_trajs,
|
||||
logprobs.reshape(-1),
|
||||
)
|
||||
)
|
||||
|
||||
# normalize reward with running variance if specified
|
||||
if self.reward_scale_running:
|
||||
reward_trajs_transpose = self.running_reward_scaler(
|
||||
reward=reward_trajs.T, first=firsts_trajs[:-1].T
|
||||
)
|
||||
reward_trajs = reward_trajs_transpose.T
|
||||
|
||||
# bootstrap value with GAE if not done - apply reward scaling with constant if specified
|
||||
obs_venv_ts = torch.from_numpy(obs_venv).float().to(self.device)
|
||||
with torch.no_grad():
|
||||
next_value = (
|
||||
self.model.critic(obs_venv_ts).reshape(1, -1).cpu().numpy()
|
||||
)
|
||||
advantages_trajs = np.zeros_like(reward_trajs)
|
||||
lastgaelam = 0
|
||||
for t in reversed(range(self.n_steps)):
|
||||
if t == self.n_steps - 1:
|
||||
nextnonterminal = 1.0 - done_venv
|
||||
nextvalues = next_value
|
||||
else:
|
||||
nextnonterminal = 1.0 - dones_trajs[t + 1]
|
||||
nextvalues = values_trajs[t + 1]
|
||||
# delta = r + gamma*V(st+1) - V(st)
|
||||
delta = (
|
||||
reward_trajs[t] * self.reward_scale_const
|
||||
+ self.gamma * nextvalues * nextnonterminal
|
||||
- values_trajs[t]
|
||||
)
|
||||
# A = delta_t + gamma*lamdba*delta_{t+1} + ...
|
||||
advantages_trajs[t] = lastgaelam = (
|
||||
delta
|
||||
+ self.gamma
|
||||
* self.gae_lambda
|
||||
* nextnonterminal
|
||||
* lastgaelam
|
||||
)
|
||||
returns_trajs = advantages_trajs + values_trajs
|
||||
|
||||
# k for environment step
|
||||
obs_k = einops.rearrange(
|
||||
torch.tensor(obs_trajs).float().to(self.device),
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
samples_k = einops.rearrange(
|
||||
torch.tensor(samples_trajs).float().to(self.device),
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
returns_k = (
|
||||
torch.tensor(returns_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
values_k = (
|
||||
torch.tensor(values_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
advantages_k = (
|
||||
torch.tensor(advantages_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
logprobs_k = (
|
||||
torch.tensor(logprobs_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
|
||||
# Update policy and critic
|
||||
total_steps = self.n_steps * self.n_envs
|
||||
inds_k = np.arange(total_steps)
|
||||
clipfracs = []
|
||||
for update_epoch in range(self.update_epochs):
|
||||
|
||||
# for each epoch, go through all data in batches
|
||||
flag_break = False
|
||||
np.random.shuffle(inds_k)
|
||||
num_batch = max(1, total_steps // self.batch_size) # skip last ones
|
||||
for batch in range(num_batch):
|
||||
start = batch * self.batch_size
|
||||
end = start + self.batch_size
|
||||
inds_b = inds_k[start:end] # b for batch
|
||||
obs_b = obs_k[inds_b]
|
||||
samples_b = samples_k[inds_b]
|
||||
returns_b = returns_k[inds_b]
|
||||
values_b = values_k[inds_b]
|
||||
advantages_b = advantages_k[inds_b]
|
||||
logprobs_b = logprobs_k[inds_b]
|
||||
|
||||
# get loss
|
||||
(
|
||||
pg_loss,
|
||||
entropy_loss,
|
||||
v_loss,
|
||||
clipfrac,
|
||||
approx_kl,
|
||||
ratio,
|
||||
bc_loss,
|
||||
std,
|
||||
) = self.model.loss(
|
||||
obs_b,
|
||||
samples_b,
|
||||
returns_b,
|
||||
values_b,
|
||||
advantages_b,
|
||||
logprobs_b,
|
||||
use_bc_loss=self.use_bc_loss,
|
||||
)
|
||||
loss = (
|
||||
pg_loss
|
||||
+ entropy_loss * self.ent_coef
|
||||
+ v_loss * self.vf_coef
|
||||
+ bc_loss * self.bc_loss_coeff
|
||||
)
|
||||
clipfracs += [clipfrac]
|
||||
|
||||
# update policy and critic
|
||||
self.actor_optimizer.zero_grad()
|
||||
self.critic_optimizer.zero_grad()
|
||||
loss.backward()
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
if self.max_grad_norm is not None:
|
||||
torch.nn.utils.clip_grad_norm_(
|
||||
self.model.actor_ft.parameters(), self.max_grad_norm
|
||||
)
|
||||
self.actor_optimizer.step()
|
||||
self.critic_optimizer.step()
|
||||
log.info(
|
||||
f"approx_kl: {approx_kl}, update_epoch: {update_epoch}, num_batch: {num_batch}"
|
||||
)
|
||||
|
||||
# Stop gradient update if KL difference reaches target
|
||||
if self.target_kl is not None and approx_kl > self.target_kl:
|
||||
flag_break = True
|
||||
break
|
||||
if flag_break:
|
||||
break
|
||||
|
||||
# Explained variation of future rewards using value function
|
||||
y_pred, y_true = values_k.cpu().numpy(), returns_k.cpu().numpy()
|
||||
var_y = np.var(y_true)
|
||||
explained_var = (
|
||||
np.nan if var_y == 0 else 1 - np.var(y_true - y_pred) / var_y
|
||||
)
|
||||
|
||||
# Plot state trajectories
|
||||
if (
|
||||
self.itr % self.render_freq == 0
|
||||
and self.n_render > 0
|
||||
and self.traj_plotter is not None
|
||||
):
|
||||
self.traj_plotter(
|
||||
obs_full_trajs=obs_full_trajs,
|
||||
n_render=self.n_render,
|
||||
max_episode_steps=self.max_episode_steps,
|
||||
render_dir=self.render_dir,
|
||||
itr=self.itr,
|
||||
)
|
||||
|
||||
# Update lr
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
self.actor_lr_scheduler.step()
|
||||
self.critic_lr_scheduler.step()
|
||||
|
||||
# Save model
|
||||
if self.itr % self.save_model_freq == 0 or self.itr == self.n_train_itr - 1:
|
||||
self.save_model()
|
||||
|
||||
# Log loss and save metrics
|
||||
run_results.append(
|
||||
{
|
||||
"itr": self.itr,
|
||||
}
|
||||
)
|
||||
if self.save_trajs:
|
||||
run_results[-1]["obs_full_trajs"] = obs_full_trajs
|
||||
run_results[-1]["obs_trajs"] = obs_trajs
|
||||
run_results[-1]["action_trajs"] = samples_trajs
|
||||
run_results[-1]["reward_trajs"] = reward_trajs
|
||||
if self.itr % self.log_freq == 0:
|
||||
time = timer()
|
||||
if eval_mode:
|
||||
log.info(
|
||||
f"eval: success rate {success_rate:8.4f} | avg episode reward {avg_episode_reward:8.4f} | avg best reward {avg_best_reward:8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"success rate - eval": success_rate,
|
||||
"avg episode reward - eval": avg_episode_reward,
|
||||
"avg best reward - eval": avg_best_reward,
|
||||
"num episode - eval": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=False,
|
||||
)
|
||||
run_results[-1]["eval_success_rate"] = success_rate
|
||||
run_results[-1]["eval_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["eval_best_reward"] = avg_best_reward
|
||||
else:
|
||||
log.info(
|
||||
f"{self.itr}: loss {loss:8.4f} | pg loss {pg_loss:8.4f} | value loss {v_loss:8.4f} | ent {-entropy_loss:8.4f} | reward {avg_episode_reward:8.4f} | t:{time:8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"loss": loss,
|
||||
"pg loss": pg_loss,
|
||||
"value loss": v_loss,
|
||||
"entropy": -entropy_loss,
|
||||
"std": std,
|
||||
"approx kl": approx_kl,
|
||||
"ratio": ratio,
|
||||
"clipfrac": np.mean(clipfracs),
|
||||
"explained variance": explained_var,
|
||||
"avg episode reward - train": avg_episode_reward,
|
||||
"num episode - train": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=True,
|
||||
)
|
||||
run_results[-1]["loss"] = loss
|
||||
run_results[-1]["pg_loss"] = pg_loss
|
||||
run_results[-1]["value_loss"] = v_loss
|
||||
run_results[-1]["entropy_loss"] = entropy_loss
|
||||
run_results[-1]["approx_kl"] = approx_kl
|
||||
run_results[-1]["ratio"] = ratio
|
||||
run_results[-1]["clip_frac"] = np.mean(clipfracs)
|
||||
run_results[-1]["explained_variance"] = explained_var
|
||||
run_results[-1]["train_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["time"] = time
|
||||
with open(self.result_path, "wb") as f:
|
||||
pickle.dump(run_results, f)
|
||||
self.itr += 1
|
||||
@@ -0,0 +1,443 @@
|
||||
"""
|
||||
PPO training for Gaussian/GMM policy with pixel observations.
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
import pickle
|
||||
import einops
|
||||
import numpy as np
|
||||
import torch
|
||||
import logging
|
||||
import wandb
|
||||
import math
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from util.timer import Timer
|
||||
from agent.finetune.train_ppo_gaussian_agent import TrainPPOGaussianAgent
|
||||
from model.common.modules import RandomShiftsAug
|
||||
|
||||
|
||||
class TrainPPOImgGaussianAgent(TrainPPOGaussianAgent):
|
||||
|
||||
def __init__(self, cfg):
|
||||
super().__init__(cfg)
|
||||
|
||||
# Image randomization
|
||||
self.augment = cfg.train.augment
|
||||
if self.augment:
|
||||
self.aug = RandomShiftsAug(pad=4)
|
||||
|
||||
# Set obs dim - we will save the different obs in batch in a dict
|
||||
shape_meta = cfg.shape_meta
|
||||
self.obs_dims = {k: shape_meta.obs[k]["shape"] for k in shape_meta.obs.keys()}
|
||||
|
||||
# Gradient accumulation to deal with large GPU RAM usage
|
||||
self.grad_accumulate = cfg.train.grad_accumulate
|
||||
|
||||
def run(self):
|
||||
|
||||
# Start training loop
|
||||
timer = Timer()
|
||||
run_results = []
|
||||
last_itr_eval = False
|
||||
done_venv = np.zeros((1, self.n_envs))
|
||||
while self.itr < self.n_train_itr:
|
||||
|
||||
# Prepare video paths for each envs --- only applies for the first set of episodes if allowing reset within iteration and each iteration has multiple episodes from one env
|
||||
options_venv = [{} for _ in range(self.n_envs)]
|
||||
if self.itr % self.render_freq == 0 and self.render_video:
|
||||
for env_ind in range(self.n_render):
|
||||
options_venv[env_ind]["video_path"] = os.path.join(
|
||||
self.render_dir, f"itr-{self.itr}_trial-{env_ind}.mp4"
|
||||
)
|
||||
|
||||
# Define train or eval - all envs restart
|
||||
eval_mode = self.itr % self.val_freq == 0 and not self.force_train
|
||||
self.model.eval() if eval_mode else self.model.train()
|
||||
last_itr_eval = eval_mode
|
||||
|
||||
# Reset env before iteration starts (1) if specified, (2) at eval mode, or (3) right after eval mode
|
||||
dones_trajs = np.empty((0, self.n_envs))
|
||||
firsts_trajs = np.zeros((self.n_steps + 1, self.n_envs))
|
||||
if self.reset_at_iteration or eval_mode or last_itr_eval:
|
||||
prev_obs_venv = self.reset_env_all(options_venv=options_venv)
|
||||
firsts_trajs[0] = 1
|
||||
else:
|
||||
firsts_trajs[0] = (
|
||||
done_venv # if done at the end of last iteration, then the envs are just reset
|
||||
)
|
||||
|
||||
# Holder
|
||||
obs_trajs = {
|
||||
k: np.empty((0, self.n_envs, self.n_cond_step, *self.obs_dims[k]))
|
||||
for k in self.obs_dims
|
||||
}
|
||||
samples_trajs = np.empty(
|
||||
(
|
||||
0,
|
||||
self.n_envs,
|
||||
self.horizon_steps,
|
||||
self.action_dim,
|
||||
)
|
||||
)
|
||||
reward_trajs = np.empty((0, self.n_envs))
|
||||
|
||||
# Collect a set of trajectories from env
|
||||
for step in range(self.n_steps):
|
||||
if step % 10 == 0:
|
||||
print(f"Processed step {step} of {self.n_steps}")
|
||||
|
||||
# Select action
|
||||
with torch.no_grad():
|
||||
cond = {
|
||||
key: torch.from_numpy(prev_obs_venv[key])
|
||||
.float()
|
||||
.to(self.device)
|
||||
for key in self.obs_dims.keys()
|
||||
} # batch each type of obs and put into dict
|
||||
samples = self.model(
|
||||
cond=cond,
|
||||
deterministic=eval_mode,
|
||||
)
|
||||
output_venv = samples.cpu().numpy()
|
||||
action_venv = output_venv[:, : self.act_steps]
|
||||
|
||||
# Apply multi-step action
|
||||
obs_venv, reward_venv, done_venv, info_venv = self.venv.step(
|
||||
action_venv
|
||||
)
|
||||
for k in obs_trajs.keys():
|
||||
obs_trajs[k] = np.vstack((obs_trajs[k], prev_obs_venv[k][None]))
|
||||
samples_trajs = np.vstack((samples_trajs, output_venv[None]))
|
||||
reward_trajs = np.vstack((reward_trajs, reward_venv[None]))
|
||||
dones_trajs = np.vstack((dones_trajs, done_venv[None]))
|
||||
firsts_trajs[step + 1] = done_venv
|
||||
prev_obs_venv = obs_venv
|
||||
|
||||
# Summarize episode reward --- this needs to be handled differently depending on whether the environment is reset after each iteration. Only count episodes that finish within the iteration.
|
||||
episodes_start_end = []
|
||||
for env_ind in range(self.n_envs):
|
||||
env_steps = np.where(firsts_trajs[:, env_ind] == 1)[0]
|
||||
for i in range(len(env_steps) - 1):
|
||||
start = env_steps[i]
|
||||
end = env_steps[i + 1]
|
||||
if end - start > 1:
|
||||
episodes_start_end.append((env_ind, start, end - 1))
|
||||
if len(episodes_start_end) > 0:
|
||||
reward_trajs_split = [
|
||||
reward_trajs[start : end + 1, env_ind]
|
||||
for env_ind, start, end in episodes_start_end
|
||||
]
|
||||
num_episode_finished = len(reward_trajs_split)
|
||||
episode_reward = np.array(
|
||||
[np.sum(reward_traj) for reward_traj in reward_trajs_split]
|
||||
)
|
||||
episode_best_reward = np.array(
|
||||
[
|
||||
np.max(reward_traj) / self.act_steps
|
||||
for reward_traj in reward_trajs_split
|
||||
]
|
||||
)
|
||||
avg_episode_reward = np.mean(episode_reward)
|
||||
avg_best_reward = np.mean(episode_best_reward)
|
||||
success_rate = np.mean(
|
||||
episode_best_reward >= self.best_reward_threshold_for_success
|
||||
)
|
||||
else:
|
||||
episode_reward = np.array([])
|
||||
num_episode_finished = 0
|
||||
avg_episode_reward = 0
|
||||
avg_best_reward = 0
|
||||
success_rate = 0
|
||||
log.info("[WARNING] No episode completed within the iteration!")
|
||||
|
||||
# Update
|
||||
if not eval_mode:
|
||||
with torch.no_grad():
|
||||
# apply image randomization
|
||||
obs_trajs["rgb"] = (
|
||||
torch.from_numpy(obs_trajs["rgb"]).float().to(self.device)
|
||||
)
|
||||
obs_trajs["state"] = (
|
||||
torch.from_numpy(obs_trajs["state"]).float().to(self.device)
|
||||
)
|
||||
if self.augment:
|
||||
rgb = einops.rearrange(
|
||||
obs_trajs["rgb"],
|
||||
"s e t c h w -> (s e t) c h w",
|
||||
)
|
||||
rgb = self.aug(rgb)
|
||||
obs_trajs["rgb"] = einops.rearrange(
|
||||
rgb,
|
||||
"(s e t) c h w -> s e t c h w",
|
||||
s=self.n_steps,
|
||||
e=self.n_envs,
|
||||
)
|
||||
|
||||
# Calculate value and logprobs - split into batches to prevent out of memory
|
||||
num_split = math.ceil(
|
||||
self.n_envs * self.n_steps / self.logprob_batch_size
|
||||
)
|
||||
obs_ts = [{} for _ in range(num_split)]
|
||||
for k in obs_trajs.keys():
|
||||
obs_k = einops.rearrange(
|
||||
obs_trajs[k],
|
||||
"s e ... -> (s e) ...",
|
||||
)
|
||||
obs_ts_k = torch.split(obs_k, self.logprob_batch_size, dim=0)
|
||||
for i, obs_t in enumerate(obs_ts_k):
|
||||
obs_ts[i][k] = obs_t
|
||||
values_trajs = np.empty((0, self.n_envs))
|
||||
for obs in obs_ts:
|
||||
values = (
|
||||
self.model.critic(obs, no_augment=True)
|
||||
.cpu()
|
||||
.numpy()
|
||||
.flatten()
|
||||
)
|
||||
values_trajs = np.vstack(
|
||||
(values_trajs, values.reshape(-1, self.n_envs))
|
||||
)
|
||||
samples_t = einops.rearrange(
|
||||
torch.from_numpy(samples_trajs).float().to(self.device),
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
samples_ts = torch.split(samples_t, self.logprob_batch_size, dim=0)
|
||||
logprobs_trajs = np.empty((0))
|
||||
for obs_t, samples_t in zip(obs_ts, samples_ts):
|
||||
logprobs = (
|
||||
self.model.get_logprobs(obs_t, samples_t)[0].cpu().numpy()
|
||||
)
|
||||
logprobs_trajs = np.concatenate(
|
||||
(
|
||||
logprobs_trajs,
|
||||
logprobs.reshape(-1),
|
||||
)
|
||||
)
|
||||
|
||||
# normalize reward with running variance if specified
|
||||
if self.reward_scale_running:
|
||||
reward_trajs_transpose = self.running_reward_scaler(
|
||||
reward=reward_trajs.T, first=firsts_trajs[:-1].T
|
||||
)
|
||||
reward_trajs = reward_trajs_transpose.T
|
||||
|
||||
# bootstrap value with GAE if not done - apply reward scaling with constant if specified
|
||||
obs_venv_ts = {
|
||||
key: torch.from_numpy(obs_venv[key]).float().to(self.device)
|
||||
for key in self.obs_dims.keys()
|
||||
}
|
||||
with torch.no_grad():
|
||||
next_value = (
|
||||
self.model.critic(obs_venv_ts, no_augment=True)
|
||||
.reshape(1, -1)
|
||||
.cpu()
|
||||
.numpy()
|
||||
)
|
||||
advantages_trajs = np.zeros_like(reward_trajs)
|
||||
lastgaelam = 0
|
||||
for t in reversed(range(self.n_steps)):
|
||||
if t == self.n_steps - 1:
|
||||
nextnonterminal = 1.0 - done_venv
|
||||
nextvalues = next_value
|
||||
else:
|
||||
nextnonterminal = 1.0 - dones_trajs[t + 1]
|
||||
nextvalues = values_trajs[t + 1]
|
||||
# delta = r + gamma*V(st+1) - V(st)
|
||||
delta = (
|
||||
reward_trajs[t] * self.reward_scale_const
|
||||
+ self.gamma * nextvalues * nextnonterminal
|
||||
- values_trajs[t]
|
||||
)
|
||||
# A = delta_t + gamma*lamdba*delta_{t+1} + ...
|
||||
advantages_trajs[t] = lastgaelam = (
|
||||
delta
|
||||
+ self.gamma
|
||||
* self.gae_lambda
|
||||
* nextnonterminal
|
||||
* lastgaelam
|
||||
)
|
||||
returns_trajs = advantages_trajs + values_trajs
|
||||
|
||||
# k for environment step
|
||||
obs_k = {
|
||||
k: einops.rearrange(
|
||||
obs_trajs[k],
|
||||
"s e ... -> (s e) ...",
|
||||
)
|
||||
for k in obs_trajs.keys()
|
||||
}
|
||||
samples_k = einops.rearrange(
|
||||
torch.tensor(samples_trajs).float().to(self.device),
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
returns_k = (
|
||||
torch.tensor(returns_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
values_k = (
|
||||
torch.tensor(values_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
advantages_k = (
|
||||
torch.tensor(advantages_trajs).float().to(self.device).reshape(-1)
|
||||
)
|
||||
logprobs_k = torch.tensor(logprobs_trajs).float().to(self.device)
|
||||
|
||||
# Update policy and critic
|
||||
total_steps = self.n_steps * self.n_envs
|
||||
inds_k = np.arange(total_steps)
|
||||
clipfracs = []
|
||||
for update_epoch in range(self.update_epochs):
|
||||
|
||||
# for each epoch, go through all data in batches
|
||||
flag_break = False
|
||||
np.random.shuffle(inds_k)
|
||||
num_batch = max(1, total_steps // self.batch_size) # skip last ones
|
||||
for batch in range(num_batch):
|
||||
start = batch * self.batch_size
|
||||
end = start + self.batch_size
|
||||
inds_b = inds_k[start:end] # b for batch
|
||||
obs_b = {k: obs_k[k][inds_b] for k in obs_k.keys()}
|
||||
samples_b = samples_k[inds_b]
|
||||
returns_b = returns_k[inds_b]
|
||||
values_b = values_k[inds_b]
|
||||
advantages_b = advantages_k[inds_b]
|
||||
logprobs_b = logprobs_k[inds_b]
|
||||
|
||||
# get loss
|
||||
(
|
||||
pg_loss,
|
||||
entropy_loss,
|
||||
v_loss,
|
||||
clipfrac,
|
||||
approx_kl,
|
||||
ratio,
|
||||
bc_loss,
|
||||
std,
|
||||
) = self.model.loss(
|
||||
obs_b,
|
||||
samples_b,
|
||||
returns_b,
|
||||
values_b,
|
||||
advantages_b,
|
||||
logprobs_b,
|
||||
use_bc_loss=self.use_bc_loss,
|
||||
)
|
||||
loss = (
|
||||
pg_loss
|
||||
+ entropy_loss * self.ent_coef
|
||||
+ v_loss * self.vf_coef
|
||||
+ bc_loss * self.bc_loss_coeff
|
||||
)
|
||||
clipfracs += [clipfrac]
|
||||
|
||||
# update policy and critic
|
||||
loss.backward()
|
||||
if (batch + 1) % self.grad_accumulate == 0:
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
if self.max_grad_norm is not None:
|
||||
torch.nn.utils.clip_grad_norm_(
|
||||
self.model.actor_ft.parameters(),
|
||||
self.max_grad_norm,
|
||||
)
|
||||
self.actor_optimizer.step()
|
||||
self.critic_optimizer.step()
|
||||
self.actor_optimizer.zero_grad()
|
||||
self.critic_optimizer.zero_grad()
|
||||
log.info(f"run grad update at batch {batch}")
|
||||
log.info(
|
||||
f"approx_kl: {approx_kl}, update_epoch: {update_epoch}, num_batch: {num_batch}"
|
||||
)
|
||||
|
||||
# Stop gradient update if KL difference reaches target
|
||||
if (
|
||||
self.target_kl is not None
|
||||
and approx_kl > self.target_kl
|
||||
and self.itr >= self.n_critic_warmup_itr
|
||||
):
|
||||
flag_break = True
|
||||
break
|
||||
if flag_break:
|
||||
break
|
||||
|
||||
# Explained variation of future rewards using value function
|
||||
y_pred, y_true = values_k.cpu().numpy(), returns_k.cpu().numpy()
|
||||
var_y = np.var(y_true)
|
||||
explained_var = (
|
||||
np.nan if var_y == 0 else 1 - np.var(y_true - y_pred) / var_y
|
||||
)
|
||||
|
||||
# Update lr
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
self.actor_lr_scheduler.step()
|
||||
self.critic_lr_scheduler.step()
|
||||
|
||||
# Save model
|
||||
if self.itr % self.save_model_freq == 0 or self.itr == self.n_train_itr - 1:
|
||||
self.save_model()
|
||||
|
||||
# Log loss and save metrics
|
||||
run_results.append(
|
||||
{
|
||||
"itr": self.itr,
|
||||
}
|
||||
)
|
||||
if self.itr % self.log_freq == 0:
|
||||
if eval_mode:
|
||||
log.info(
|
||||
f"eval: success rate {success_rate:8.4f} | avg episode reward {avg_episode_reward:8.4f} | avg best reward {avg_best_reward:8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"success rate - eval": success_rate,
|
||||
"avg episode reward - eval": avg_episode_reward,
|
||||
"avg best reward - eval": avg_best_reward,
|
||||
"num episode - eval": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=False,
|
||||
)
|
||||
run_results[-1]["eval_success_rate"] = success_rate
|
||||
run_results[-1]["eval_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["eval_best_reward"] = avg_best_reward
|
||||
else:
|
||||
log.info(
|
||||
f"{self.itr}: loss {loss:8.4f} | pg loss {pg_loss:8.4f} | value loss {v_loss:8.4f} | bc loss {bc_loss:8.4f} | reward {avg_episode_reward:8.4f} | t:{timer():8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"loss": loss,
|
||||
"pg loss": pg_loss,
|
||||
"value loss": v_loss,
|
||||
"bc loss": bc_loss,
|
||||
"std": std,
|
||||
"approx kl": approx_kl,
|
||||
"ratio": ratio,
|
||||
"clipfrac": np.mean(clipfracs),
|
||||
"explained variance": explained_var,
|
||||
"avg episode reward - train": avg_episode_reward,
|
||||
"num episode - train": num_episode_finished,
|
||||
"actor lr": self.actor_optimizer.param_groups[0]["lr"],
|
||||
"critic lr": self.critic_optimizer.param_groups[0][
|
||||
"lr"
|
||||
],
|
||||
},
|
||||
step=self.itr,
|
||||
commit=True,
|
||||
)
|
||||
run_results[-1]["loss"] = loss
|
||||
run_results[-1]["pg_loss"] = pg_loss
|
||||
run_results[-1]["value_loss"] = v_loss
|
||||
run_results[-1]["bc_loss"] = bc_loss
|
||||
run_results[-1]["std"] = std
|
||||
run_results[-1]["approx_kl"] = approx_kl
|
||||
run_results[-1]["ratio"] = ratio
|
||||
run_results[-1]["clip_frac"] = np.mean(clipfracs)
|
||||
run_results[-1]["explained_variance"] = explained_var
|
||||
run_results[-1]["train_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["time"] = timer()
|
||||
with open(self.result_path, "wb") as f:
|
||||
pickle.dump(run_results, f)
|
||||
self.itr += 1
|
||||
@@ -0,0 +1,316 @@
|
||||
"""
|
||||
QSM (Q-Score Matching) for diffusion policy.
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
import pickle
|
||||
import einops
|
||||
import numpy as np
|
||||
import torch
|
||||
import logging
|
||||
import wandb
|
||||
from copy import deepcopy
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from util.timer import Timer
|
||||
from collections import deque
|
||||
from agent.finetune.train_agent import TrainAgent
|
||||
from util.scheduler import CosineAnnealingWarmupRestarts
|
||||
|
||||
|
||||
class TrainQSMDiffusionAgent(TrainAgent):
|
||||
|
||||
def __init__(self, cfg):
|
||||
super().__init__(cfg)
|
||||
|
||||
# note the discount factor gamma here is applied to reward every act_steps, instead of every env step
|
||||
self.gamma = cfg.train.gamma
|
||||
|
||||
# Wwarm up period for critic before actor updates
|
||||
self.n_critic_warmup_itr = cfg.train.n_critic_warmup_itr
|
||||
|
||||
# Optimizer
|
||||
self.actor_optimizer = torch.optim.AdamW(
|
||||
self.model.actor.parameters(),
|
||||
lr=cfg.train.actor_lr,
|
||||
weight_decay=cfg.train.actor_weight_decay,
|
||||
)
|
||||
self.actor_lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.actor_optimizer,
|
||||
first_cycle_steps=cfg.train.actor_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.actor_lr,
|
||||
min_lr=cfg.train.actor_lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.actor_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
self.critic_optimizer = torch.optim.AdamW(
|
||||
self.model.critic_q.parameters(),
|
||||
lr=cfg.train.critic_lr,
|
||||
weight_decay=cfg.train.critic_weight_decay,
|
||||
)
|
||||
self.critic_lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.critic_optimizer,
|
||||
first_cycle_steps=cfg.train.critic_lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.critic_lr,
|
||||
min_lr=cfg.train.critic_lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.critic_lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
|
||||
# Buffer size
|
||||
self.buffer_size = cfg.train.buffer_size
|
||||
|
||||
# Scaling reward
|
||||
self.scale_reward_factor = cfg.train.scale_reward_factor
|
||||
|
||||
# Updates
|
||||
self.replay_ratio = cfg.train.replay_ratio
|
||||
self.critic_tau = cfg.train.critic_tau
|
||||
self.q_grad_coeff = cfg.train.q_grad_coeff
|
||||
|
||||
def run(self):
|
||||
|
||||
# make a FIFO replay buffer for obs, action, and reward
|
||||
obs_buffer = deque(maxlen=self.buffer_size)
|
||||
action_buffer = deque(maxlen=self.buffer_size)
|
||||
next_obs_buffer = deque(maxlen=self.buffer_size)
|
||||
reward_buffer = deque(maxlen=self.buffer_size)
|
||||
done_buffer = deque(maxlen=self.buffer_size)
|
||||
first_buffer = deque(maxlen=self.buffer_size)
|
||||
|
||||
# Start training loop
|
||||
timer = Timer()
|
||||
run_results = []
|
||||
done_venv = np.zeros((1, self.n_envs))
|
||||
while self.itr < self.n_train_itr:
|
||||
|
||||
# Prepare video paths for each envs --- only applies for the first set of episodes if allowing reset within iteration and each iteration has multiple episodes from one env
|
||||
options_venv = [{} for _ in range(self.n_envs)]
|
||||
if self.itr % self.render_freq == 0 and self.render_video:
|
||||
for env_ind in range(self.n_render):
|
||||
options_venv[env_ind]["video_path"] = os.path.join(
|
||||
self.render_dir, f"itr-{self.itr}_trial-{env_ind}.mp4"
|
||||
)
|
||||
|
||||
# Define train or eval - all envs restart
|
||||
eval_mode = self.itr % self.val_freq == 0 and not self.force_train
|
||||
self.model.eval() if eval_mode else self.model.train()
|
||||
firsts_trajs = np.zeros((self.n_steps + 1, self.n_envs))
|
||||
|
||||
# Reset env at the beginning of an iteration
|
||||
if self.reset_at_iteration or eval_mode or last_itr_eval:
|
||||
prev_obs_venv = self.reset_env_all(options_venv=options_venv)
|
||||
firsts_trajs[0] = 1
|
||||
else:
|
||||
firsts_trajs[0] = (
|
||||
done_venv # if done at the end of last iteration, then the envs are just reset
|
||||
)
|
||||
last_itr_eval = eval_mode
|
||||
reward_trajs = np.empty((0, self.n_envs))
|
||||
|
||||
# Collect a set of trajectories from env
|
||||
for step in range(self.n_steps):
|
||||
if step % 10 == 0:
|
||||
print(f"Processed step {step} of {self.n_steps}")
|
||||
|
||||
# Select action
|
||||
with torch.no_grad():
|
||||
samples = (
|
||||
self.model(
|
||||
cond=torch.from_numpy(prev_obs_venv)
|
||||
.float()
|
||||
.to(self.device),
|
||||
deterministic=eval_mode,
|
||||
)
|
||||
.cpu()
|
||||
.numpy()
|
||||
) # n_env x horizon x act
|
||||
action_venv = samples[:, : self.act_steps]
|
||||
|
||||
# Apply multi-step action
|
||||
obs_venv, reward_venv, done_venv, info_venv = self.venv.step(
|
||||
action_venv
|
||||
)
|
||||
reward_trajs = np.vstack((reward_trajs, reward_venv[None]))
|
||||
|
||||
# add to buffer
|
||||
obs_buffer.append(prev_obs_venv)
|
||||
action_buffer.append(action_venv)
|
||||
next_obs_buffer.append(obs_venv)
|
||||
reward_buffer.append(reward_venv * self.scale_reward_factor)
|
||||
done_buffer.append(done_venv)
|
||||
first_buffer.append(firsts_trajs[step])
|
||||
|
||||
firsts_trajs[step + 1] = done_venv
|
||||
prev_obs_venv = obs_venv
|
||||
|
||||
# Summarize episode reward --- this needs to be handled differently depending on whether the environment is reset after each iteration. Only count episodes that finish within the iteration.
|
||||
episodes_start_end = []
|
||||
for env_ind in range(self.n_envs):
|
||||
env_steps = np.where(firsts_trajs[:, env_ind] == 1)[0]
|
||||
for i in range(len(env_steps) - 1):
|
||||
start = env_steps[i]
|
||||
end = env_steps[i + 1]
|
||||
if end - start > 1:
|
||||
episodes_start_end.append((env_ind, start, end - 1))
|
||||
if len(episodes_start_end) > 0:
|
||||
reward_trajs_split = [
|
||||
reward_trajs[start : end + 1, env_ind]
|
||||
for env_ind, start, end in episodes_start_end
|
||||
]
|
||||
num_episode_finished = len(reward_trajs_split)
|
||||
episode_reward = np.array(
|
||||
[np.sum(reward_traj) for reward_traj in reward_trajs_split]
|
||||
)
|
||||
episode_best_reward = np.array(
|
||||
[
|
||||
np.max(reward_traj) / self.act_steps
|
||||
for reward_traj in reward_trajs_split
|
||||
]
|
||||
)
|
||||
avg_episode_reward = np.mean(episode_reward)
|
||||
avg_best_reward = np.mean(episode_best_reward)
|
||||
success_rate = np.mean(
|
||||
episode_best_reward >= self.best_reward_threshold_for_success
|
||||
)
|
||||
else:
|
||||
episode_reward = np.array([])
|
||||
num_episode_finished = 0
|
||||
avg_episode_reward = 0
|
||||
avg_best_reward = 0
|
||||
success_rate = 0
|
||||
log.info("[WARNING] No episode completed within the iteration!")
|
||||
|
||||
# Update
|
||||
if not eval_mode:
|
||||
|
||||
obs_trajs = np.array(deepcopy(obs_buffer))
|
||||
action_trajs = np.array(deepcopy(action_buffer))
|
||||
next_obs_trajs = np.array(deepcopy(next_obs_buffer))
|
||||
reward_trajs = np.array(deepcopy(reward_buffer))
|
||||
done_trajs = np.array(deepcopy(done_buffer))
|
||||
first_trajs = np.array(deepcopy(first_buffer))
|
||||
|
||||
# flatten
|
||||
obs_trajs = einops.rearrange(
|
||||
obs_trajs,
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
next_obs_trajs = einops.rearrange(
|
||||
next_obs_trajs,
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
action_trajs = einops.rearrange(
|
||||
action_trajs,
|
||||
"s e h d -> (s e) h d",
|
||||
)
|
||||
reward_trajs = reward_trajs.reshape(-1)
|
||||
done_trajs = done_trajs.reshape(-1)
|
||||
first_trajs = first_trajs.reshape(-1)
|
||||
|
||||
num_batch = int(
|
||||
self.n_steps * self.n_envs / self.batch_size * self.replay_ratio
|
||||
)
|
||||
|
||||
for _ in range(num_batch):
|
||||
|
||||
# Sample batch
|
||||
inds = np.random.choice(len(obs_trajs), self.batch_size)
|
||||
obs_b = torch.from_numpy(obs_trajs[inds]).float().to(self.device)
|
||||
next_obs_b = (
|
||||
torch.from_numpy(next_obs_trajs[inds]).float().to(self.device)
|
||||
)
|
||||
actions_b = (
|
||||
torch.from_numpy(action_trajs[inds]).float().to(self.device)
|
||||
)
|
||||
reward_b = (
|
||||
torch.from_numpy(reward_trajs[inds]).float().to(self.device)
|
||||
)
|
||||
done_b = torch.from_numpy(done_trajs[inds]).float().to(self.device)
|
||||
|
||||
# update critic q function
|
||||
critic_loss = self.model.loss_critic(
|
||||
obs_b, next_obs_b, actions_b, reward_b, done_b, self.gamma
|
||||
)
|
||||
self.critic_optimizer.zero_grad()
|
||||
critic_loss.backward()
|
||||
self.critic_optimizer.step()
|
||||
|
||||
# update target q function
|
||||
self.model.update_target_critic(self.critic_tau)
|
||||
|
||||
loss_critic = critic_loss.detach()
|
||||
|
||||
# Update policy with collected trajectories
|
||||
loss = self.model.loss_actor(
|
||||
obs_b,
|
||||
actions_b,
|
||||
self.q_grad_coeff,
|
||||
)
|
||||
self.actor_optimizer.zero_grad()
|
||||
loss.backward()
|
||||
if self.itr >= self.n_critic_warmup_itr:
|
||||
if self.max_grad_norm is not None:
|
||||
torch.nn.utils.clip_grad_norm_(
|
||||
self.model.actor.parameters(), self.max_grad_norm
|
||||
)
|
||||
self.actor_optimizer.step()
|
||||
|
||||
# Update lr
|
||||
self.actor_lr_scheduler.step()
|
||||
self.critic_lr_scheduler.step()
|
||||
|
||||
# Save model
|
||||
if self.itr % self.save_model_freq == 0 or self.itr == self.n_train_itr - 1:
|
||||
self.save_model()
|
||||
|
||||
# Log loss and save metrics
|
||||
run_results.append(
|
||||
{
|
||||
"itr": self.itr,
|
||||
}
|
||||
)
|
||||
if self.itr % self.log_freq == 0:
|
||||
if eval_mode:
|
||||
log.info(
|
||||
f"eval: success rate {success_rate:8.4f} | avg episode reward {avg_episode_reward:8.4f} | avg best reward {avg_best_reward:8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"success rate - eval": success_rate,
|
||||
"avg episode reward - eval": avg_episode_reward,
|
||||
"avg best reward - eval": avg_best_reward,
|
||||
"num episode - eval": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=False,
|
||||
)
|
||||
run_results[-1]["eval_success_rate"] = success_rate
|
||||
run_results[-1]["eval_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["eval_best_reward"] = avg_best_reward
|
||||
else:
|
||||
log.info(
|
||||
f"{self.itr}: loss {loss:8.4f} | reward {avg_episode_reward:8.4f} |t:{timer():8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"loss": loss,
|
||||
"loss - critic": loss_critic,
|
||||
"avg episode reward - train": avg_episode_reward,
|
||||
"num episode - train": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=True,
|
||||
)
|
||||
run_results[-1]["loss"] = loss
|
||||
run_results[-1]["loss_critic"] = loss_critic
|
||||
run_results[-1]["train_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["time"] = timer()
|
||||
with open(self.result_path, "wb") as f:
|
||||
pickle.dump(run_results, f)
|
||||
self.itr += 1
|
||||
@@ -0,0 +1,301 @@
|
||||
"""
|
||||
Reward-weighted regression (RWR) for diffusion policy.
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
import pickle
|
||||
import numpy as np
|
||||
import torch
|
||||
import logging
|
||||
import wandb
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from util.timer import Timer
|
||||
from agent.finetune.train_agent import TrainAgent
|
||||
from util.scheduler import CosineAnnealingWarmupRestarts
|
||||
|
||||
|
||||
class TrainRWRDiffusionAgent(TrainAgent):
|
||||
|
||||
def __init__(self, cfg):
|
||||
super().__init__(cfg)
|
||||
|
||||
# note the discount factor gamma here is applied to reward every act_steps, instead of every env step
|
||||
self.gamma = cfg.train.gamma
|
||||
|
||||
# Build optimizer
|
||||
self.optimizer = torch.optim.AdamW(
|
||||
self.model.parameters(),
|
||||
lr=cfg.train.lr,
|
||||
weight_decay=cfg.train.weight_decay,
|
||||
)
|
||||
self.lr_scheduler = CosineAnnealingWarmupRestarts(
|
||||
self.optimizer,
|
||||
first_cycle_steps=cfg.train.lr_scheduler.first_cycle_steps,
|
||||
cycle_mult=1.0,
|
||||
max_lr=cfg.train.lr,
|
||||
min_lr=cfg.train.lr_scheduler.min_lr,
|
||||
warmup_steps=cfg.train.lr_scheduler.warmup_steps,
|
||||
gamma=1.0,
|
||||
)
|
||||
|
||||
# Reward exponential
|
||||
self.beta = cfg.train.beta
|
||||
|
||||
# Max weight for AWR
|
||||
self.max_reward_weight = cfg.train.max_reward_weight
|
||||
|
||||
# Updates
|
||||
self.update_epochs = cfg.train.update_epochs
|
||||
|
||||
def run(self):
|
||||
|
||||
# Start training loop
|
||||
timer = Timer()
|
||||
run_results = []
|
||||
done_venv = np.zeros((1, self.n_envs))
|
||||
while self.itr < self.n_train_itr:
|
||||
|
||||
# Prepare video paths for each envs --- only applies for the first set of episodes if allowing reset within iteration and each iteration has multiple episodes from one env
|
||||
options_venv = [{} for _ in range(self.n_envs)]
|
||||
if self.itr % self.render_freq == 0 and self.render_video:
|
||||
for env_ind in range(self.n_render):
|
||||
options_venv[env_ind]["video_path"] = os.path.join(
|
||||
self.render_dir, f"itr-{self.itr}_trial-{env_ind}.mp4"
|
||||
)
|
||||
|
||||
# Define train or eval - all envs restart
|
||||
eval_mode = self.itr % self.val_freq == 0 and not self.force_train
|
||||
self.model.eval() if eval_mode else self.model.train()
|
||||
firsts_trajs = np.zeros((self.n_steps + 1, self.n_envs))
|
||||
|
||||
# Reset env at the beginning of an iteration
|
||||
if self.reset_at_iteration or eval_mode or last_itr_eval:
|
||||
prev_obs_venv = self.reset_env_all(options_venv=options_venv)
|
||||
firsts_trajs[0] = 1
|
||||
else:
|
||||
firsts_trajs[0] = (
|
||||
done_venv # if done at the end of last iteration, then the envs are just reset
|
||||
)
|
||||
last_itr_eval = eval_mode
|
||||
reward_trajs = np.empty((0, self.n_envs))
|
||||
|
||||
# Holders
|
||||
obs_trajs = np.empty((0, self.n_envs, self.n_cond_step, self.obs_dim))
|
||||
samples_trajs = np.empty(
|
||||
(
|
||||
0,
|
||||
self.n_envs,
|
||||
self.horizon_steps,
|
||||
self.action_dim,
|
||||
)
|
||||
)
|
||||
|
||||
# Collect a set of trajectories from env
|
||||
for step in range(self.n_steps):
|
||||
if step % 10 == 0:
|
||||
print(f"Processed step {step} of {self.n_steps}")
|
||||
|
||||
# Select action
|
||||
with torch.no_grad():
|
||||
samples = (
|
||||
self.model(
|
||||
cond=torch.from_numpy(prev_obs_venv)
|
||||
.float()
|
||||
.to(self.device),
|
||||
deterministic=eval_mode,
|
||||
)
|
||||
.cpu()
|
||||
.numpy()
|
||||
) # n_env x horizon x act
|
||||
action_venv = samples[:, : self.act_steps]
|
||||
obs_trajs = np.vstack((obs_trajs, prev_obs_venv[None]))
|
||||
samples_trajs = np.vstack((samples_trajs, samples[None]))
|
||||
|
||||
# Apply multi-step action
|
||||
obs_venv, reward_venv, done_venv, info_venv = self.venv.step(
|
||||
action_venv
|
||||
)
|
||||
reward_trajs = np.vstack((reward_trajs, reward_venv[None]))
|
||||
firsts_trajs[step + 1] = done_venv
|
||||
prev_obs_venv = obs_venv
|
||||
|
||||
# Summarize episode reward --- this needs to be handled differently depending on whether the environment is reset after each iteration. Only count episodes that finish within the iteration.
|
||||
episodes_start_end = []
|
||||
for env_ind in range(self.n_envs):
|
||||
env_steps = np.where(firsts_trajs[:, env_ind] == 1)[0]
|
||||
for i in range(len(env_steps) - 1):
|
||||
start = env_steps[i]
|
||||
end = env_steps[i + 1]
|
||||
if end - start > 1:
|
||||
episodes_start_end.append((env_ind, start, end - 1))
|
||||
if len(episodes_start_end) > 0:
|
||||
# Compute transitions for completed trajectories
|
||||
obs_trajs_split = [
|
||||
obs_trajs[start : end + 1, env_ind]
|
||||
for env_ind, start, end in episodes_start_end
|
||||
]
|
||||
samples_trajs_split = [
|
||||
samples_trajs[start : end + 1, env_ind]
|
||||
for env_ind, start, end in episodes_start_end
|
||||
]
|
||||
reward_trajs_split = [
|
||||
reward_trajs[start : end + 1, env_ind]
|
||||
for env_ind, start, end in episodes_start_end
|
||||
]
|
||||
num_episode_finished = len(reward_trajs_split)
|
||||
|
||||
# Compute episode returns
|
||||
discounted_reward_trajs_split = [
|
||||
[
|
||||
self.gamma**t * r
|
||||
for t, r in zip(
|
||||
list(range(end - start + 1)),
|
||||
reward_trajs[start : end + 1, env_ind],
|
||||
)
|
||||
]
|
||||
for env_ind, start, end in episodes_start_end
|
||||
]
|
||||
returns_trajs_split = [
|
||||
np.cumsum(y[::-1])[::-1] for y in discounted_reward_trajs_split
|
||||
]
|
||||
returns_trajs_split = np.concatenate(returns_trajs_split)
|
||||
episode_reward = np.array(
|
||||
[np.sum(reward_traj) for reward_traj in reward_trajs_split]
|
||||
)
|
||||
episode_best_reward = np.array(
|
||||
[
|
||||
np.max(reward_traj) / self.act_steps
|
||||
for reward_traj in reward_trajs_split
|
||||
]
|
||||
)
|
||||
avg_episode_reward = np.mean(episode_reward)
|
||||
avg_best_reward = np.mean(episode_best_reward)
|
||||
success_rate = np.mean(
|
||||
episode_best_reward >= self.best_reward_threshold_for_success
|
||||
)
|
||||
else:
|
||||
episode_reward = np.array([])
|
||||
num_episode_finished = 0
|
||||
avg_episode_reward = 0
|
||||
avg_best_reward = 0
|
||||
success_rate = 0
|
||||
log.info("[WARNING] No episode completed within the iteration!")
|
||||
|
||||
# Update
|
||||
if not eval_mode:
|
||||
|
||||
# Tensorize data and put them to device
|
||||
# k for environment step
|
||||
obs_k = (
|
||||
torch.tensor(np.concatenate(obs_trajs_split))
|
||||
.float()
|
||||
.to(self.device)
|
||||
)
|
||||
|
||||
samples_k = (
|
||||
torch.tensor(np.concatenate(samples_trajs_split))
|
||||
.float()
|
||||
.to(self.device)
|
||||
)
|
||||
|
||||
# Normalize reward
|
||||
returns_trajs_split = (
|
||||
returns_trajs_split - np.mean(returns_trajs_split)
|
||||
) / (returns_trajs_split.std() + 1e-3)
|
||||
|
||||
rewards_k = (
|
||||
torch.tensor(returns_trajs_split)
|
||||
.float()
|
||||
.to(self.device)
|
||||
.reshape(-1)
|
||||
)
|
||||
|
||||
rewards_k_scaled = torch.exp(self.beta * rewards_k)
|
||||
rewards_k_scaled.clamp_(max=self.max_reward_weight)
|
||||
|
||||
# rewards_k_scaled = rewards_k_scaled / rewards_k_scaled.mean()
|
||||
|
||||
# Update policy and critic
|
||||
total_steps = len(rewards_k_scaled)
|
||||
inds_k = np.arange(total_steps)
|
||||
for _ in range(self.update_epochs):
|
||||
|
||||
# for each epoch, go through all data in batches
|
||||
np.random.shuffle(inds_k)
|
||||
num_batch = max(1, total_steps // self.batch_size) # skip last ones
|
||||
for batch in range(num_batch):
|
||||
start = batch * self.batch_size
|
||||
end = start + self.batch_size
|
||||
inds_b = inds_k[start:end] # b for batch
|
||||
obs_b = obs_k[inds_b]
|
||||
samples_b = samples_k[inds_b]
|
||||
rewards_b = rewards_k_scaled[inds_b]
|
||||
|
||||
# Update policy with collected trajectories
|
||||
loss = self.model.loss(
|
||||
samples_b,
|
||||
obs_b,
|
||||
rewards_b,
|
||||
)
|
||||
self.optimizer.zero_grad()
|
||||
loss.backward()
|
||||
if self.max_grad_norm is not None:
|
||||
torch.nn.utils.clip_grad_norm_(
|
||||
self.model.parameters(), self.max_grad_norm
|
||||
)
|
||||
self.optimizer.step()
|
||||
|
||||
# Update lr
|
||||
self.lr_scheduler.step()
|
||||
|
||||
# Save model
|
||||
if self.itr % self.save_model_freq == 0 or self.itr == self.n_train_itr - 1:
|
||||
self.save_model()
|
||||
|
||||
# Log loss and save metrics
|
||||
run_results.append(
|
||||
{
|
||||
"itr": self.itr,
|
||||
}
|
||||
)
|
||||
if self.itr % self.log_freq == 0:
|
||||
if eval_mode:
|
||||
log.info(
|
||||
f"eval: success rate {success_rate:8.4f} | avg episode reward {avg_episode_reward:8.4f} | avg best reward {avg_best_reward:8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"success rate - eval": success_rate,
|
||||
"avg episode reward - eval": avg_episode_reward,
|
||||
"avg best reward - eval": avg_best_reward,
|
||||
"num episode - eval": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=False,
|
||||
)
|
||||
run_results[-1]["eval_success_rate"] = success_rate
|
||||
run_results[-1]["eval_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["eval_best_reward"] = avg_best_reward
|
||||
else:
|
||||
log.info(
|
||||
f"{self.itr}: loss {loss:8.4f} | reward {avg_episode_reward:8.4f} |t:{timer():8.4f}"
|
||||
)
|
||||
if self.use_wandb:
|
||||
wandb.log(
|
||||
{
|
||||
"loss": loss,
|
||||
"avg episode reward - train": avg_episode_reward,
|
||||
"num episode - train": num_episode_finished,
|
||||
},
|
||||
step=self.itr,
|
||||
commit=True,
|
||||
)
|
||||
run_results[-1]["loss"] = loss
|
||||
run_results[-1]["train_episode_reward"] = avg_episode_reward
|
||||
run_results[-1]["time"] = timer()
|
||||
with open(self.result_path, "wb") as f:
|
||||
pickle.dump(run_results, f)
|
||||
self.itr += 1
|
||||
Reference in New Issue
Block a user