* 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
+23 -12
View File
@@ -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,
+1 -1
View File
@@ -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
View File
@@ -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,