* Sampling over both env and denoising steps in DPPO updates (#13)

* sample one from each chain

* full random sampling

* Add Proficient Human (PH) Configs and Pipeline (#16)

* fix missing cfg

* add ph config

* fix how terminated flags are added to buffer in ibrl

* add ph config

* offline calql for 1M gradient updates

* bug fix: number of calql online gradient steps is the number of new transitions collected

* add sample config for DPPO with ta=1

* Sampling over both env and denoising steps in DPPO updates (#13)

* sample one from each chain

* full random sampling

* fix diffusion loss when predicting initial noise

* fix dppo inds

* fix typo

* remove print statement

---------

Co-authored-by: Justin M. Lidard <jlidard@neuronic.cs.princeton.edu>
Co-authored-by: allenzren <allen.ren@princeton.edu>

* update robomimic configs

* better calql formulation

* optimize calql and ibrl training

* optimize data transfer in ppo agents

* add kitchen configs

* re-organize config folders, rerun calql and rlpd

* add scratch gym locomotion configs

* add kitchen installation dependencies

* use truncated for termination in furniture env

* update furniture and gym configs

* update README and dependencies with kitchen

* add url for new data and checkpoints

* update demo RL configs

* update batch sizes for furniture unet configs

* raise error about dropout in residual mlp

* fix observation bug in bc loss

---------

Co-authored-by: Justin Lidard <60638575+jlidard@users.noreply.github.com>
Co-authored-by: Justin M. Lidard <jlidard@neuronic.cs.princeton.edu>
This commit is contained in:
Allen Z. Ren
2024-10-30 19:58:06 -04:00
committed by GitHub
co-authored by Justin M. Lidard Justin Lidard
parent 7b10df690d
commit dc8e0c9edc
126 changed files with 4614 additions and 553 deletions
+16 -21
View File
@@ -22,7 +22,6 @@ from model.diffusion.diffusion_vpg import VPGDiffusion
class PPODiffusion(VPGDiffusion):
def __init__(
self,
gamma_denoising: float,
@@ -58,7 +57,9 @@ class PPODiffusion(VPGDiffusion):
def loss(
self,
obs,
chains,
chains_prev,
chains_next,
denoising_inds,
returns,
oldvalues,
advantages,
@@ -81,9 +82,11 @@ class PPODiffusion(VPGDiffusion):
reward_horizon: action horizon that backpropagates gradient
"""
# Get new logprobs for denoising steps from T-1 to 0 - entropy is fixed fod diffusion
newlogprobs, eta = self.get_logprobs(
newlogprobs, eta = self.get_logprobs_subsample(
obs,
chains,
chains_prev,
chains_next,
denoising_inds,
get_ent=True,
)
entropy_loss = -eta.mean()
@@ -92,7 +95,7 @@ class PPODiffusion(VPGDiffusion):
# only backpropagate through the earlier steps (e.g., ones actually executed in the environment)
newlogprobs = newlogprobs[:, :reward_horizon, :]
oldlogprobs = oldlogprobs[:, :, :reward_horizon, :]
oldlogprobs = oldlogprobs[:, :reward_horizon, :]
# Get the logprobs - batch over B and denoising steps
newlogprobs = newlogprobs.mean(dim=(-1, -2)).view(-1)
@@ -106,9 +109,7 @@ class PPODiffusion(VPGDiffusion):
# Get counterfactual teacher actions
samples = self.forward(
cond=obs.float()
.unsqueeze(1)
.to(self.device), # B x horizon=1 x obs_dim
cond=obs,
deterministic=False,
return_chain=True,
use_base_policy=True,
@@ -116,7 +117,7 @@ class PPODiffusion(VPGDiffusion):
# Get logprobs of teacher actions under this policy
bc_logprobs = self.get_logprobs(
obs,
samples.chains, # n_env x denoising x horizon x act
samples.chains,
get_ent=False,
use_base_policy=False,
)
@@ -133,14 +134,13 @@ class PPODiffusion(VPGDiffusion):
advantage_max = torch.quantile(advantages, self.clip_advantage_upper_quantile)
advantages = advantages.clamp(min=advantage_min, max=advantage_max)
# repeat advantages for denoising steps and horizon steps
advantages = advantages.repeat_interleave(self.ft_denoising_steps)
# denoising discount
discount = torch.tensor(
[self.gamma_denoising**i for i in reversed(range(self.ft_denoising_steps))]
[
self.gamma_denoising ** (self.ft_denoising_steps - i - 1)
for i in denoising_inds
]
).to(self.device)
discount = discount.repeat(len(advantages) // self.ft_denoising_steps)
advantages *= discount
# get ratio
@@ -148,9 +148,7 @@ class PPODiffusion(VPGDiffusion):
ratio = logratio.exp()
# exponentially interpolate between the base and the current clipping value over denoising steps and repeat
t = torch.arange(self.ft_denoising_steps).float().to(self.device) / (
self.ft_denoising_steps - 1
) # 0 to 1
t = (denoising_inds.float() / (self.ft_denoising_steps - 1)).to(self.device)
if self.ft_denoising_steps > 1:
clip_ploss_coef = self.clip_ploss_coef_base + (
self.clip_ploss_coef - self.clip_ploss_coef_base
@@ -158,10 +156,7 @@ class PPODiffusion(VPGDiffusion):
math.exp(self.clip_ploss_coef_rate) - 1
)
else:
clip_ploss_coef = torch.tensor([self.clip_ploss_coef]).to(self.device)
clip_ploss_coef = clip_ploss_coef.repeat(
len(advantages) // self.ft_denoising_steps
)
clip_ploss_coef = t
# get kl difference and whether value clipped
with torch.no_grad():
+65
View File
@@ -395,6 +395,71 @@ class VPGDiffusion(DiffusionModel):
return log_prob, eta
return log_prob
def get_logprobs_subsample(
self,
cond,
chains_prev,
chains_next,
denoising_inds,
get_ent: bool = False,
use_base_policy: bool = False,
):
"""
Calculating the logprobs of random samples of denoised chains.
Args:
cond: dict with key state/rgb; more recent obs at the end
state: (B, To, Do)
rgb: (B, To, C, H, W)
chains: (B, K+1, Ta, Da)
get_ent: flag for returning entropy
use_base_policy: flag for using base policy
Returns:
logprobs: (B, Ta, Da)
entropy (if get_ent=True): (B, Ta)
denoising_indices: (B, )
"""
# Sample t for batch dim, keep it 1-dim
if self.use_ddim:
t_single = self.ddim_t[-self.ft_denoising_steps :]
else:
t_single = torch.arange(
start=self.ft_denoising_steps - 1,
end=-1,
step=-1,
device=self.device,
)
# 4,3,2,1,0,4,3,2,1,0,...,4,3,2,1,0
t_all = t_single[denoising_inds]
if self.use_ddim:
ddim_indices_single = torch.arange(
start=self.ddim_steps - self.ft_denoising_steps,
end=self.ddim_steps,
device=self.device,
) # only used for DDIM
ddim_indices = ddim_indices_single[denoising_inds]
else:
ddim_indices = None
# Forward pass with previous chains
next_mean, logvar, eta = self.p_mean_var(
chains_prev,
t_all,
cond=cond,
index=ddim_indices,
use_base_policy=use_base_policy,
)
std = torch.exp(0.5 * logvar)
std = torch.clip(std, min=self.min_logprob_denoising_std)
dist = Normal(next_mean, std)
# Get logprobs with gaussian
log_prob = dist.log_prob(chains_next)
if get_ent:
return log_prob, eta
return log_prob
def loss(self, cond, chains, reward):
"""
REINFORCE loss. Not used right now.