v0.6 (#18)
* 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:
co-authored by
Justin M. Lidard
Justin Lidard
parent
7b10df690d
commit
dc8e0c9edc
@@ -96,6 +96,7 @@ class ResidualMLP(nn.Module):
|
||||
out_activation_type="Identity",
|
||||
use_layernorm=False,
|
||||
use_layernorm_final=False,
|
||||
dropout=0,
|
||||
):
|
||||
super(ResidualMLP, self).__init__()
|
||||
hidden_dim = dim_list[1]
|
||||
@@ -108,6 +109,7 @@ class ResidualMLP(nn.Module):
|
||||
hidden_dim=hidden_dim,
|
||||
activation_type=activation_type,
|
||||
use_layernorm=use_layernorm,
|
||||
dropout=dropout,
|
||||
)
|
||||
for _ in range(1, num_hidden_layers, 2)
|
||||
]
|
||||
@@ -129,6 +131,7 @@ class TwoLayerPreActivationResNetLinear(nn.Module):
|
||||
hidden_dim,
|
||||
activation_type="Mish",
|
||||
use_layernorm=False,
|
||||
dropout=0,
|
||||
):
|
||||
super().__init__()
|
||||
self.l1 = nn.Linear(hidden_dim, hidden_dim)
|
||||
@@ -137,6 +140,8 @@ class TwoLayerPreActivationResNetLinear(nn.Module):
|
||||
if use_layernorm:
|
||||
self.norm1 = nn.LayerNorm(hidden_dim, eps=1e-06)
|
||||
self.norm2 = nn.LayerNorm(hidden_dim, eps=1e-06)
|
||||
if dropout > 0:
|
||||
raise NotImplementedError("Dropout not implemented for residual MLP!")
|
||||
|
||||
def forward(self, x):
|
||||
x_input = x
|
||||
|
||||
@@ -212,6 +212,7 @@ class Gaussian_MLP(nn.Module):
|
||||
out_activation_type=activation_type,
|
||||
use_layernorm=use_layernorm,
|
||||
use_layernorm_final=use_layernorm,
|
||||
dropout=dropout,
|
||||
)
|
||||
self.mlp_mean = MLP(
|
||||
mlp_dims[-1:] + [output_dim],
|
||||
@@ -233,9 +234,7 @@ class Gaussian_MLP(nn.Module):
|
||||
if learn_fixed_std:
|
||||
# initialize to fixed_std
|
||||
self.logvar = torch.nn.Parameter(
|
||||
torch.log(
|
||||
torch.tensor([fixed_std**2 for _ in range(action_dim)])
|
||||
),
|
||||
torch.log(torch.tensor([fixed_std**2 for _ in range(action_dim)])),
|
||||
requires_grad=True,
|
||||
)
|
||||
self.logvar_min = torch.nn.Parameter(
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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.
|
||||
|
||||
+23
-12
@@ -63,7 +63,6 @@ class CalQL_Gaussian(GaussianModel):
|
||||
returns,
|
||||
terminated,
|
||||
gamma,
|
||||
alpha,
|
||||
):
|
||||
B = len(actions)
|
||||
|
||||
@@ -71,17 +70,17 @@ class CalQL_Gaussian(GaussianModel):
|
||||
q_data1, q_data2 = self.critic(obs, actions)
|
||||
with torch.no_grad():
|
||||
# repeat for action samples
|
||||
next_obs["state"] = next_obs["state"].repeat_interleave(
|
||||
next_obs_repeated = {"state": next_obs["state"].repeat_interleave(
|
||||
self.cql_n_actions, dim=0
|
||||
)
|
||||
)}
|
||||
|
||||
# Get the next actions and logprobs
|
||||
next_actions, next_logprobs = self.forward(
|
||||
next_obs,
|
||||
next_obs_repeated,
|
||||
deterministic=False,
|
||||
get_logprob=True,
|
||||
)
|
||||
next_q1, next_q2 = self.target_critic(next_obs, next_actions)
|
||||
next_q1, next_q2 = self.target_critic(next_obs_repeated, next_actions)
|
||||
next_q = torch.min(next_q1, next_q2)
|
||||
|
||||
# Reshape the next_q to match the number of samples
|
||||
@@ -96,9 +95,6 @@ class CalQL_Gaussian(GaussianModel):
|
||||
# Get the target Q values
|
||||
target_q = rewards + gamma * (1 - terminated) * next_q
|
||||
|
||||
# Subtract the entropy bonus
|
||||
target_q = target_q - alpha * next_logprobs
|
||||
|
||||
# TD loss
|
||||
td_loss_1 = nn.functional.mse_loss(q_data1, target_q)
|
||||
td_loss_2 = nn.functional.mse_loss(q_data2, target_q)
|
||||
@@ -111,6 +107,12 @@ class CalQL_Gaussian(GaussianModel):
|
||||
reparameterize=False,
|
||||
get_logprob=True,
|
||||
) # no gradient
|
||||
pi_next_actions, log_pi_next = self.forward(
|
||||
next_obs,
|
||||
deterministic=False,
|
||||
reparameterize=False,
|
||||
get_logprob=True,
|
||||
) # no gradient
|
||||
|
||||
# Random action Q values
|
||||
n_random_actions = random_actions.shape[1]
|
||||
@@ -130,17 +132,26 @@ class CalQL_Gaussian(GaussianModel):
|
||||
|
||||
# Policy action Q values
|
||||
q_pi_1, q_pi_2 = self.critic(obs, pi_actions)
|
||||
q_pi_1 = q_pi_1 - log_pi
|
||||
q_pi_2 = q_pi_2 - log_pi
|
||||
q_pi_next_1, q_pi_next_2 = self.critic(next_obs, pi_next_actions)
|
||||
|
||||
# Ensure calibration w.r.t. value function estimate
|
||||
q_pi_1 = torch.max(q_pi_1, returns)[:, None] # (B, 1)
|
||||
q_pi_2 = torch.max(q_pi_2, returns)[:, None] # (B, 1)
|
||||
cat_q_1 = torch.cat([q_rand_1, q_pi_1], dim=-1) # (B, num_samples+1)
|
||||
q_pi_next_1 = torch.max(q_pi_next_1, returns)[:, None] # (B, 1)
|
||||
q_pi_next_2 = torch.max(q_pi_next_2, returns)[:, None] # (B, 1)
|
||||
|
||||
# cql_importance_sample
|
||||
q_pi_1 = q_pi_1 - log_pi
|
||||
q_pi_2 = q_pi_2 - log_pi
|
||||
q_pi_next_1 = q_pi_next_1 - log_pi_next
|
||||
q_pi_next_2 = q_pi_next_2 - log_pi_next
|
||||
cat_q_1 = torch.cat([q_rand_1, q_pi_1, q_pi_next_1], dim=-1) # (B, num_samples+1)
|
||||
cql_qf1_ood = torch.logsumexp(cat_q_1, dim=-1) # max over num_samples
|
||||
cat_q_2 = torch.cat([q_rand_2, q_pi_2], dim=-1) # (B, num_samples+1)
|
||||
cat_q_2 = torch.cat([q_rand_2, q_pi_2, q_pi_next_2], dim=-1) # (B, num_samples+1)
|
||||
cql_qf2_ood = torch.logsumexp(cat_q_2, dim=-1) # sum over num_samples
|
||||
|
||||
# skip cal_lagrange since the paper shows cql_target_action_gap not used in kitchen
|
||||
|
||||
# Subtract the log likelihood of the data
|
||||
cql_qf1_diff = torch.clamp(
|
||||
cql_qf1_ood - q_data1,
|
||||
|
||||
@@ -20,7 +20,7 @@ class IBRL_Gaussian(GaussianModel):
|
||||
critic,
|
||||
n_critics,
|
||||
soft_action_sample=False,
|
||||
soft_action_sample_beta=0.1,
|
||||
soft_action_sample_beta=10,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(network=actor, **kwargs)
|
||||
|
||||
+17
-19
@@ -63,6 +63,23 @@ class PPO_Gaussian(VPG_Gaussian):
|
||||
oldlogprobs = oldlogprobs.clamp(min=-5, max=2)
|
||||
entropy_loss = -entropy
|
||||
|
||||
bc_loss = 0.0
|
||||
if use_bc_loss:
|
||||
# See Eqn. 2 of https://arxiv.org/pdf/2403.03949.pdf
|
||||
# Give a reward for maximizing probability of teacher policy's action with current policy.
|
||||
# Actions are chosen along trajectory induced by current policy.
|
||||
|
||||
# Get counterfactual teacher actions
|
||||
samples = self.forward(
|
||||
cond=obs,
|
||||
deterministic=False,
|
||||
use_base_policy=True,
|
||||
)
|
||||
# Get logprobs of teacher actions under this policy
|
||||
bc_logprobs, _, _ = self.get_logprobs(obs, samples, use_base_policy=False)
|
||||
bc_logprobs = bc_logprobs.clamp(min=-5, max=2)
|
||||
bc_loss = -bc_logprobs.mean()
|
||||
|
||||
# get ratio
|
||||
logratio = newlogprobs - oldlogprobs
|
||||
ratio = logratio.exp()
|
||||
@@ -99,25 +116,6 @@ class PPO_Gaussian(VPG_Gaussian):
|
||||
v_loss = 0.5 * v_loss_max.mean()
|
||||
else:
|
||||
v_loss = 0.5 * ((newvalues - returns) ** 2).mean()
|
||||
|
||||
bc_loss = 0.0
|
||||
if use_bc_loss:
|
||||
# See Eqn. 2 of https://arxiv.org/pdf/2403.03949.pdf
|
||||
# Give a reward for maximizing probability of teacher policy's action with current policy.
|
||||
# Actions are chosen along trajectory induced by current policy.
|
||||
|
||||
# Get counterfactual teacher actions
|
||||
samples = self.forward(
|
||||
cond=obs.float()
|
||||
.unsqueeze(1)
|
||||
.to(self.device), # B x horizon=1 x obs_dim
|
||||
deterministic=False,
|
||||
use_base_policy=True,
|
||||
)
|
||||
# Get logprobs of teacher actions under this policy
|
||||
bc_logprobs, _, _ = self.get_logprobs(obs, samples, use_base_policy=False)
|
||||
bc_logprobs = bc_logprobs.clamp(min=-5, max=2)
|
||||
bc_loss = -bc_logprobs.mean()
|
||||
return (
|
||||
pg_loss,
|
||||
entropy_loss,
|
||||
|
||||
Reference in New Issue
Block a user