squash commits

This commit is contained in:
allenzren
2024-09-11 21:09:17 -04:00
parent 8ce0aa1485
commit 2ddf63b8f5
200 changed files with 1240 additions and 1186 deletions
+8 -59
View File
@@ -8,7 +8,7 @@ Annotated DDIM/DDPM: https://nn.labml.ai/diffusion/stable_diffusion/sampler/ddpm
"""
from typing import Optional, Union
from typing import Union
import logging
import torch
from torch import nn
@@ -17,13 +17,12 @@ import torch.nn.functional as F
log = logging.getLogger(__name__)
from model.diffusion.sampling import (
make_timesteps,
extract,
cosine_beta_schedule,
)
from collections import namedtuple
Sample = namedtuple("Sample", "trajectories values chains")
Sample = namedtuple("Sample", "trajectories chains")
class DiffusionModel(nn.Module):
@@ -34,9 +33,7 @@ class DiffusionModel(nn.Module):
horizon_steps,
obs_dim,
action_dim,
transition_dim,
network_path=None,
cond_steps=1,
device="cuda:0",
# DDPM parameters
denoising_steps=100,
@@ -53,11 +50,9 @@ class DiffusionModel(nn.Module):
self.horizon_steps = horizon_steps
self.obs_dim = obs_dim
self.action_dim = action_dim
self.transition_dim = transition_dim
self.denoising_steps = int(denoising_steps)
self.denoised_clip_value = denoised_clip_value
self.predict_epsilon = predict_epsilon
self.cond_steps = cond_steps
self.use_ddim = use_ddim
self.ddim_steps = ddim_steps
@@ -216,52 +211,11 @@ class DiffusionModel(nn.Module):
@torch.no_grad()
def forward(
self,
cond: Optional[torch.Tensor],
cond,
return_chain=True,
**kwargs,
):
"""
Forward sampling through denoising steps.
Args:
cond: (batch_size, horizon, transition_dim)
return_chain: whether to return the chain of samples or only the final denoised sample
Return:
Sample: namedtuple with fields:
trajectories: (batch_size, horizon_steps, transition_dim)
values: (batch_size, )
chain: (batch_size, denoising_steps + 1, horizon_steps, transition_dim)
"""
device = self.betas.device
if isinstance(cond, dict):
B = cond[list(cond.keys())[0]].shape[0]
else:
B = cond.shape[0]
cond = cond[:, : self.cond_steps].reshape(B, -1)
shape = (B, self.horizon_steps, self.transition_dim)
# Loop
x = torch.randn(shape, device=device)
chain = [x] if return_chain else None
if self.use_ddim:
t_all = self.ddim_t
else:
t_all = list(reversed(range(self.denoising_steps)))
for i, t in enumerate(t_all):
t_b = make_timesteps(B, t, device)
index_b = make_timesteps(B, i, device)
mu, logvar = self.p_mean_var(x=x, t=t_b, cond=cond, index=index_b)
std = torch.exp(0.5 * logvar)
# no noise when t == 0
noise = torch.randn_like(x)
noise[t == 0] = 0
x = mu + std * noise
if return_chain:
chain.append(x)
if return_chain:
chain = torch.stack(chain, dim=1)
values = torch.zeros(len(x), device=x.device) # not considering the value for now
return Sample(x, values, chain)
raise NotImplementedError
# ---------- Supervised training ----------#
@@ -275,23 +229,18 @@ class DiffusionModel(nn.Module):
def p_losses(
self,
x_start,
obs_cond: Union[dict, torch.Tensor],
cond: Union[dict, torch.Tensor],
t,
):
"""
If predicting epsilon: E_{t, x0, ε} [||ε - ε_θ(√α̅ₜx0 + √(1-α̅ₜ)ε, t)||²
Args:
x_start: (batch_size, horizon_steps, transition_dim)
obs_cond: dict with keys as step and value as observation
x_start: (batch_size, horizon_steps, action_dim)
cond: dict with keys as step and value as observation
t: batch of integers
"""
device = x_start.device
B = x_start.shape[0]
if isinstance(obs_cond[0], dict):
cond = obs_cond[0] # keep the dictionary and the network will extract img and prio
else:
cond = obs_cond[0].reshape(B, -1)
# Forward process
noise = torch.randn_like(x_start, device=device)
+9 -12
View File
@@ -40,19 +40,19 @@ class DIPODiffusion(DiffusionModel):
# Whether to clamp sampled action between [-1, 1]
self.clamp_action = clamp_action
# ---------- RL training ----------#
def loss_critic(self, obs, next_obs, actions, rewards, dones, gamma):
# get current Q-function
actions_flat = torch.flatten(actions, start_dim=-2)
current_q1, current_q2 = self.critic(obs, actions_flat)
current_q1, current_q2 = self.critic(obs, actions)
# get next Q-function
next_actions = self.forward(
cond=next_obs,
deterministic=False,
) # in DiffusionModel, forward() has no gradient, which is desired here.
next_actions_flat = torch.flatten(next_actions, start_dim=-2)
next_q1, next_q2 = self.critic(next_obs, next_actions_flat)
) # forward() has no gradient, which is desired here.
next_q1, next_q2 = self.critic(next_obs, next_actions)
next_q = torch.min(next_q1, next_q2)
# terminal state mask
@@ -73,6 +73,8 @@ class DIPODiffusion(DiffusionModel):
return loss_critic
# ---------- Sampling ----------#``
# override
@torch.no_grad()
def forward(
@@ -81,15 +83,10 @@ class DIPODiffusion(DiffusionModel):
deterministic=False,
):
device = self.betas.device
B = cond.shape[0]
if isinstance(cond, dict):
raise NotImplementedError("Not implemented for images")
else:
B = cond.shape[0]
cond = cond[:, : self.cond_steps]
B = len(cond["state"])
# Loop
x = torch.randn((B, self.horizon_steps, self.transition_dim), device=device)
x = torch.randn((B, self.horizon_steps, self.action_dim), device=device)
t_all = list(reversed(range(self.denoising_steps)))
for i, t in enumerate(t_all):
t_b = make_timesteps(B, t, device)
+12 -20
View File
@@ -41,19 +41,19 @@ class DQLDiffusion(DiffusionModel):
# Whether to clamp sampled action between [-1, 1]
self.clamp_action = clamp_action
# ---------- RL training ----------#
def loss_critic(self, obs, next_obs, actions, rewards, dones, gamma):
# get current Q-function
actions_flat = torch.flatten(actions, start_dim=-2)
current_q1, current_q2 = self.critic(obs, actions_flat)
current_q1, current_q2 = self.critic(obs, actions)
# get next Q-function
next_actions = self.forward(
cond=next_obs,
deterministic=False,
) # in DiffusionModel, forward() has no gradient, which is desired here.
next_actions_flat = torch.flatten(next_actions, start_dim=-2)
next_q1, next_q2 = self.critic(next_obs, next_actions_flat)
) # forward() has no gradient, which is desired here.
next_q1, next_q2 = self.critic(next_obs, next_actions)
next_q = torch.min(next_q1, next_q2)
# terminal state mask
@@ -75,7 +75,7 @@ class DQLDiffusion(DiffusionModel):
return loss_critic
def loss_actor(self, obs, actions, q1, q2, eta):
bc_loss = self.loss(actions, {0: obs})
bc_loss = self.loss(actions, obs)
if np.random.uniform() > 0.5:
q_loss = -q1.mean() / q2.abs().mean().detach()
else:
@@ -83,6 +83,8 @@ class DQLDiffusion(DiffusionModel):
actor_loss = bc_loss + eta * q_loss
return actor_loss
# ---------- Sampling ----------#``
# override
@torch.no_grad()
def forward(
@@ -91,15 +93,10 @@ class DQLDiffusion(DiffusionModel):
deterministic=False,
):
device = self.betas.device
B = cond.shape[0]
if isinstance(cond, dict):
raise NotImplementedError("Not implemented for images")
else:
B = cond.shape[0]
cond = cond[:, : self.cond_steps]
B = len(cond["state"])
# Loop
x = torch.randn((B, self.horizon_steps, self.transition_dim), device=device)
x = torch.randn((B, self.horizon_steps, self.action_dim), device=device)
t_all = list(reversed(range(self.denoising_steps)))
for i, t in enumerate(t_all):
t_b = make_timesteps(B, t, device)
@@ -136,15 +133,10 @@ class DQLDiffusion(DiffusionModel):
Differentiable forward pass used in actor training.
"""
device = self.betas.device
B = cond.shape[0]
if isinstance(cond, dict):
raise NotImplementedError("Not implemented for images")
else:
B = cond.shape[0]
cond = cond[:, : self.cond_steps]
B = len(cond["state"])
# Loop
x = torch.randn((B, self.horizon_steps, self.transition_dim), device=device)
x = torch.randn((B, self.horizon_steps, self.action_dim), device=device)
t_all = list(reversed(range(self.denoising_steps)))
for i, t in enumerate(t_all):
t_b = make_timesteps(B, t, device)
+43 -48
View File
@@ -42,12 +42,13 @@ class IDQLDiffusion(RWRDiffusion):
# assign actor
self.actor = self.network
# ---------- RL training ----------#
def compute_advantages(self, obs, actions):
# get current Q-function
actions_flat = torch.flatten(actions, start_dim=-2)
with torch.no_grad(): # no gradients for q-function when we update value function
current_q1, current_q2 = self.target_q(obs, actions_flat)
# get current Q-function, stop gradient
with torch.no_grad():
current_q1, current_q2 = self.target_q(obs, actions)
q = torch.min(current_q1, current_q2)
# get the current V-function
@@ -59,7 +60,6 @@ class IDQLDiffusion(RWRDiffusion):
return adv
def loss_critic_v(self, obs, actions):
adv = self.compute_advantages(obs, actions)
# get the value loss
@@ -70,11 +70,10 @@ class IDQLDiffusion(RWRDiffusion):
def loss_critic_q(self, obs, next_obs, actions, rewards, dones, gamma):
# get current Q-function
actions_flat = torch.flatten(actions, start_dim=-2)
current_q1, current_q2 = self.critic_q(obs, actions_flat)
current_q1, current_q2 = self.critic_q(obs, actions)
# get the next V-function
with torch.no_grad(): # no gradients for value function when we update q function
# get the next V-function, stop gradient
with torch.no_grad():
next_v = self.critic_v(next_obs)
# terminal state mask
@@ -98,8 +97,35 @@ class IDQLDiffusion(RWRDiffusion):
def update_target_critic(self, tau):
soft_update(self.target_q, self.critic_q, tau)
# override
def p_losses(
self,
x_start,
cond,
t,
):
device = x_start.device
# Forward process
noise = torch.randn_like(x_start, device=device)
x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise)
# Predict
x_recon = self.network(x_noisy, t, cond=cond)
# Loss with mask
if self.predict_epsilon:
loss = F.mse_loss(x_recon, noise, reduction="none")
else:
loss = F.mse_loss(x_recon, x_start, reduction="none")
loss = einops.reduce(loss, "b h d -> b", "mean")
return loss.mean()
# ---------- Sampling ----------#``
# override
@torch.no_grad()
def forward( # override
def forward(
self,
cond,
deterministic=False,
@@ -107,23 +133,23 @@ class IDQLDiffusion(RWRDiffusion):
critic_hyperparam=0.7, # sampling weight for implicit policy
use_expectile_exploration=True,
):
"""assume state-only, no rgb in cond"""
# repeat obs num_sample times along dim 0
cond_shape_repeat_dims = tuple(1 for _ in cond.shape)
B, T, D = cond.shape
cond_shape_repeat_dims = tuple(1 for _ in cond["state"].shape)
B, T, D = cond["state"].shape
S = num_sample
cond_repeat = cond[None].repeat(num_sample, *cond_shape_repeat_dims)
cond_repeat = cond["state"][None].repeat(num_sample, *cond_shape_repeat_dims)
cond_repeat = cond_repeat.view(-1, T, D) # [B*S, T, D]
# for eval, use less noisy samples --- there is still DDPM noise, but final action uses small min_sampling_std
samples = super(IDQLDiffusion, self).forward(
cond_repeat,
{"state": cond_repeat},
deterministic=deterministic,
)
_, H, A = samples.shape
# get current Q-function
actions_flat = torch.flatten(samples, start_dim=-2)
current_q1, current_q2 = self.target_q(cond_repeat, actions_flat)
current_q1, current_q2 = self.target_q({"state": cond_repeat}, samples)
q = torch.min(current_q1, current_q2)
q = q.view(S, B)
@@ -141,7 +167,7 @@ class IDQLDiffusion(RWRDiffusion):
# Sample as an implicit policy for exploration
else:
# get the current value function for probabilistic exploration
current_v = self.critic_v(cond_repeat)
current_v = self.critic_v({"state": cond_repeat})
v = current_v.view(S, B)
adv = q - v
@@ -164,34 +190,3 @@ class IDQLDiffusion(RWRDiffusion):
# squeeze dummy dimension
samples = samples_best[0]
return samples
# override
def p_losses(
self,
x_start,
obs_cond,
t,
):
device = x_start.device
B, T, D = x_start.shape
# handle different ways of passing observation
if isinstance(obs_cond[0], dict):
cond = obs_cond[0]
else:
cond = obs_cond.reshape(B, -1)
# Forward process
noise = torch.randn_like(x_start, device=device)
x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise)
# Predict
x_recon = self.network(x_noisy, t, cond=cond)
# Loss with mask
if self.predict_epsilon:
loss = F.mse_loss(x_recon, noise, reduction="none")
else:
loss = F.mse_loss(x_recon, x_start, reduction="none")
loss = einops.reduce(loss, "b h d -> b", "mean")
return loss.mean()
+15 -4
View File
@@ -1,6 +1,15 @@
"""
DPPO: Diffusion Policy Policy Optimization.
K: number of denoising steps
To: observation sequence length
Ta: action chunk size
Do: observation dimension
Da: action dimension
C: image channels
H, W: image height and width
"""
from typing import Optional
@@ -60,13 +69,15 @@ class PPODiffusion(VPGDiffusion):
"""
PPO loss
obs: (B, obs_step, obs_dim)
chains: (B, num_denoising_step+1, horizon_step, action_dim)
obs: 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)
returns: (B, )
values: (B, )
advantages: (B,)
oldlogprobs: (B, num_denoising_step, horizon_step, action_dim)
use_bc_loss: add BC regularization loss
oldlogprobs: (B, K, Ta, Da)
use_bc_loss: whether to add BC regularization loss
reward_horizon: action horizon that backpropagates gradient
"""
# Get new logprobs for denoising steps from T-1 to 0 - entropy is fixed fod diffusion
+18 -3
View File
@@ -3,6 +3,11 @@ Diffusion policy gradient with exact likelihood estimation.
Based on score_sde_pytorch https://github.com/yang-song/score_sde_pytorch
To: observation sequence length
Ta: action chunk size
Do: observation dimension
Da: action dimension
"""
import torch
@@ -52,13 +57,12 @@ class PPOExactDiffusion(PPODiffusion):
num_epsilon=sde_num_epsilon,
)
def get_exact_logprobs(self, obs, samples):
def get_exact_logprobs(self, cond, samples):
"""Use torchdiffeq
samples: B x horizon x transition_dim
samples: (B x Ta x Da)
"""
# TODO: image input
cond = obs.reshape(-1, self.obs_dim)
return self.likelihood_fn(
self.actor,
self.actor_ft,
@@ -79,6 +83,17 @@ class PPOExactDiffusion(PPODiffusion):
use_bc_loss=False,
**kwargs,
):
"""
PPO loss
obs: dict with key state/rgb; more recent obs at the end
state: (B, To, Do)
samples: (B, Ta, Da)
returns: (B, )
values: (B, )
advantages: (B,)
oldlogprobs: (B, )
"""
# Get new logprobs for final x
newlogprobs = self.get_exact_logprobs(obs, samples)
newlogprobs = newlogprobs.clamp(min=-5, max=2)
+12 -15
View File
@@ -39,11 +39,12 @@ class QSMDiffusion(RWRDiffusion):
# assign actor
self.actor = self.network
# ---------- RL training ----------#
def loss_actor(self, obs, actions, q_grad_coeff):
x_start = actions
device = x_start.device
B, T, D = x_start.shape
cond = obs.reshape(B, -1)
B = len(x_start)
# Forward process
noise = torch.randn_like(x_start, device=device)
@@ -53,39 +54,35 @@ class QSMDiffusion(RWRDiffusion):
x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise)
# get current value for noisy actions as the code does --- the algorthm block in the paper is wrong, it says using a_t, the final denoised action
x_noisy_flat = torch.flatten(x_noisy, start_dim=-2)
x_noisy_flat.requires_grad_(True)
current_q1, current_q2 = self.critic_q(obs, x_noisy_flat)
# x_noisy_flat = torch.flatten(x_noisy, start_dim=-2)
x_noisy.requires_grad_(True)
current_q1, current_q2 = self.critic_q(obs, x_noisy)
# Compute dQ/da|a=noise_actions
gradient_q1 = torch.autograd.grad(current_q1.sum(), x_noisy_flat)[0]
gradient_q2 = torch.autograd.grad(current_q2.sum(), x_noisy_flat)[0]
gradient_q1 = torch.autograd.grad(current_q1.sum(), x_noisy)[0]
gradient_q2 = torch.autograd.grad(current_q2.sum(), x_noisy)[0]
gradient_q = torch.stack((gradient_q1, gradient_q2), 0).mean(0).detach()
# Predict noise from noisy actions
x_recon = self.network(x_noisy, t, cond=cond)
x_recon = torch.flatten(x_recon, start_dim=-2)
x_recon = self.network(x_noisy, t, cond=obs)
# Loss with mask - align predicted noise with critic gradient of noisy actions
# Note: the gradient of mu wrt. epsilon has a negative sign
loss = F.mse_loss(-x_recon, q_grad_coeff * gradient_q, reduction="none").mean()
return loss
def loss_critic(self, obs, next_obs, actions, rewards, dones, gamma):
# get current Q-function
actions_flat = torch.flatten(actions, start_dim=-2)
current_q1, current_q2 = self.critic_q(obs, actions_flat)
current_q1, current_q2 = self.critic_q(obs, actions)
# get next Q-function - with noise, same as QSM https://github.com/Alescontrela/score_matching_rl/blob/f02a21969b17e322eb229ceb2b0f5a9111b1b968/jaxrl5/agents/score_matching/score_matching_learner.py#L193
next_actions = self.forward(
cond=next_obs,
deterministic=False,
) # in DiffusionModel, forward() has no gradient, which is desired here.
next_actions_flat = torch.flatten(next_actions, start_dim=-2)
) # forward() has no gradient, which is desired here.
with torch.no_grad():
next_q1, next_q2 = self.target_q(next_obs, next_actions_flat)
next_q1, next_q2 = self.target_q(next_obs, next_actions)
next_q = torch.min(next_q1, next_q2)
# terminal state mask
+4 -16
View File
@@ -43,18 +43,11 @@ class RWRDiffusion(DiffusionModel):
def p_losses(
self,
x_start,
obs_cond,
cond,
rewards,
t,
):
device = x_start.device
B, T, D = x_start.shape
# handle different ways of passing observation
if isinstance(obs_cond[0], dict):
cond = obs_cond[0]
else:
cond = obs_cond.reshape(B, -1)
# Forward process
noise = torch.randn_like(x_start, device=device)
@@ -79,7 +72,7 @@ class RWRDiffusion(DiffusionModel):
self,
x,
t,
cond=None,
cond,
):
noise = self.network(x, t, cond=cond)
@@ -116,15 +109,10 @@ class RWRDiffusion(DiffusionModel):
deterministic=False,
):
device = self.betas.device
B = cond.shape[0]
if isinstance(cond, dict):
raise NotImplementedError("Not implemented for images")
else:
B = cond.shape[0]
cond = cond[:, : self.cond_steps]
B = len(cond["state"])
# Loop
x = torch.randn((B, self.horizon_steps, self.transition_dim), device=device)
x = torch.randn((B, self.horizon_steps, self.action_dim), device=device)
t_all = list(reversed(range(self.denoising_steps)))
for i, t in enumerate(t_all):
t_b = make_timesteps(B, t, device)
+65 -82
View File
@@ -1,14 +1,20 @@
"""
Policy gradient with diffusion policy.
Policy gradient with diffusion policy. VPG: vanilla policy gradient
VPG: vanilla policy gradient
K: number of denoising steps
To: observation sequence length
Ta: action chunk size
Do: observation dimension
Da: action dimension
C: image channels
H, W: image height and width
"""
import copy
import torch
import logging
import einops
log = logging.getLogger(__name__)
import torch.nn.functional as F
@@ -106,7 +112,11 @@ class VPGDiffusion(DiffusionModel):
# ---------- Sampling ----------#
def step(self):
"""Update min_sampling_denoising_std annealing and fine-tuning denoising steps annealing. Both not used currently"""
"""
Anneal min_sampling_denoising_std and fine-tuning denoising steps
Current configs do not apply annealing
"""
# anneal min_sampling_denoising_std
if type(self.min_sampling_denoising_std) is not float:
self.min_sampling_denoising_std.step()
@@ -142,7 +152,7 @@ class VPGDiffusion(DiffusionModel):
self,
x,
t,
cond=None,
cond,
index=None,
use_base_policy=False,
deterministic=False,
@@ -160,12 +170,8 @@ class VPGDiffusion(DiffusionModel):
# overwrite noise for fine-tuning steps
if len(ft_indices) > 0:
if cond is not None:
if isinstance(cond, dict):
cond = {key: cond[key][ft_indices] for key in cond}
else:
cond = cond[ft_indices]
noise_ft = actor(x[ft_indices], t[ft_indices], cond=cond)
cond_ft = {key: cond[key][ft_indices] for key in cond}
noise_ft = actor(x[ft_indices], t[ft_indices], cond=cond_ft)
noise[ft_indices] = noise_ft
# Predict x_0
@@ -208,7 +214,8 @@ class VPGDiffusion(DiffusionModel):
if deterministic:
etas = torch.zeros((x.shape[0], 1, 1)).to(x.device)
else:
etas = self.eta(cond).unsqueeze(1) # B x 1 x (transition_dim or 1)
# TODO: eta cond
etas = self.eta(cond).unsqueeze(1) # B x 1 x (Da or 1)
sigma = (
etas
* ((1 - alpha_prev) / (1 - alpha) * (1 - alpha / alpha_prev)) ** 0.5
@@ -242,30 +249,26 @@ class VPGDiffusion(DiffusionModel):
Forward pass for sampling actions.
Args:
cond: (batch_size, obs_step, obs_dim)
deterministic: whether to sample deterministically
return_chain: whether to return the chain of samples
use_base_policy: whether to use the base policy instead
cond: dict with key state/rgb; more recent obs at the end
state: (B, To, Do)
rgb: (B, To, C, H, W)
deterministic: If true, then std=0 with DDIM, or with DDPM, use normal schedule (instead of clipping at a higher value)
return_chain: whether to return the entire chain of denoised actions
use_base_policy: whether to use the frozen pre-trained policy instead
Return:
Sample: namedtuple with fields:
trajectories: (batch_size, horizon_steps, transition_dim)
values: (batch_size, )
chain: (batch_size, denoising_steps + 1, horizon_steps, transition_dim)
trajectories: (B, Ta, Da)
chain: (B, K + 1, Ta, Da)
"""
device = self.betas.device
if isinstance(cond, dict):
B = cond["state"].shape[0]
cond["state"] = cond["state"][:, : self.cond_steps]
cond["rgb"] = cond["rgb"][:, : self.cond_steps]
else:
B = cond.shape[0]
cond = cond[:, : self.cond_steps]
sample_data = cond["state"] if "state" in cond else cond["rgb"]
B = len(sample_data)
# Get updated minimum sampling denoising std
min_sampling_denoising_std = self.get_min_sampling_denoising_std()
# Loop
x = torch.randn((B, self.horizon_steps, self.transition_dim), device=device)
x = torch.randn((B, self.horizon_steps, self.action_dim), device=device)
if self.use_ddim:
t_all = self.ddim_t
else:
@@ -317,17 +320,16 @@ class VPGDiffusion(DiffusionModel):
self.ddim_steps - self.ft_denoising_steps - 1
):
chain.append(x)
values = torch.zeros(len(x), device=x.device)
if return_chain:
chain = torch.stack(chain, dim=1)
return Sample(x, values, chain)
return Sample(x, chain)
# ---------- RL training ----------#
def get_logprobs(
self,
obs,
cond,
chains,
get_ent: bool = False,
use_base_policy: bool = False,
@@ -336,35 +338,25 @@ class VPGDiffusion(DiffusionModel):
Calculating the logprobs of the entire chain of denoised actions.
Args:
obs: (B, obs_step, obs_dim)
chains: (B, num_denoising_step+1, horizon_step, action_dim)
get_ent: flag for returning entropy
use_base_policy: flag for using base policy
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 x num_denoising_steps, horizon_step, action_dim)
entropy (if get_ent=True): (B x num_denoising_steps, horizon_step)
logprobs: (B x K, Ta, Da)
entropy (if get_ent=True): (B x K, Ta)
"""
# Repeat obs conditioning for denoising_steps
if isinstance(obs, dict):
obs = {
key: obs[key]
.unsqueeze(1)
.repeat(1, self.ft_denoising_steps, *(1,) * (obs[key].ndim - 1))
for key in obs
}
else:
obs = einops.repeat(obs, "b h d -> b t h d", t=self.ft_denoising_steps)
# flatten the first two dimensions
if isinstance(obs, dict):
cond = obs
for key in cond:
cond[key] = einops.rearrange(cond[key], "b t ... -> (b t) ...")
cond[key] = cond[key][:, : self.cond_steps]
else:
cond = einops.rearrange(obs, "b t h d -> (b t) h d")
cond = cond[:, : self.cond_steps]
# Repeat cond for denoising_steps, flatten batch and time dimensions
cond = {
key: cond[key]
.unsqueeze(1)
.repeat(1, self.ft_denoising_steps, *(1,) * (cond[key].ndim - 1))
.flatten(start_dim=0, end_dim=1)
for key in cond
} # less memory usage than einops?
# Repeat t for batch dim, keep it 1-dim
if self.use_ddim:
@@ -393,8 +385,8 @@ class VPGDiffusion(DiffusionModel):
chains_next = chains[:, 1:]
# Flatten first two dimensions
chains_prev = chains_prev.reshape(-1, self.horizon_steps, self.transition_dim)
chains_next = chains_next.reshape(-1, self.horizon_steps, self.transition_dim)
chains_prev = chains_prev.reshape(-1, self.horizon_steps, self.action_dim)
chains_next = chains_next.reshape(-1, self.horizon_steps, self.action_dim)
# Forward pass with previous chains
next_mean, logvar, eta = self.p_mean_var(
@@ -414,44 +406,35 @@ class VPGDiffusion(DiffusionModel):
return log_prob, eta
return log_prob
def loss(self, obs, chains, reward):
def loss(self, cond, chains, reward):
"""
REINFORCE loss. Not used right now.
Args:
obs: (n_steps, n_envs, obs_dim)
chains: (n_steps, n_envs, num_denoising_step+1, horizon_step, action_dim)
reward (to go): (n_steps, n_envs)
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)
reward (to go): (b,)
"""
if torch.is_tensor(reward):
assert not reward.requires_grad
n_steps, n_envs, _ = obs.shape
# Flatten first two dimensions
obs = einops.rearrange(obs, "s e d -> (s e) d")
chains = einops.rearrange(chains, "s e t h d -> (s e) t h d")
reward = reward.reshape(-1)
# Get advantage
with torch.no_grad():
value = self.critic(obs).squeeze()
value = self.critic(cond).squeeze()
advantage = reward - value
# Get logprobs for denoising steps from T-1 to 0
logprobs, eta = self.get_logprobs(obs, chains, get_ent=True)
# (n_steps x n_envs x denoising_steps) x horizon_steps x (obs_dim+action_dim)
logprobs, eta = self.get_logprobs(cond, chains, get_ent=True)
# (n_steps x n_envs x K) x Ta x (Do+Da)
# Ignore obs dimension, and then sum over action dimension
logprobs = logprobs[:, :, : self.action_dim].sum(-1)
# -> (n_steps x n_envs x denoising_steps) x horizon_steps
# -> (n_steps x n_envs x K) x Ta
# -> (n_steps x n_envs) x K x Ta
logprobs = logprobs.reshape((-1, self.denoising_steps, self.horizon_steps))
# -> (n_steps x n_envs) x denoising_steps x horizon_steps
logprobs = logprobs.reshape(
(n_steps * n_envs, self.denoising_steps, self.horizon_steps)
)
# Sum/avg over denoising steps
logprobs = logprobs.mean(-2) # -> (n_steps x n_envs) x horizon_steps
logprobs = logprobs.mean(-2) # -> (n_steps x n_envs) x Ta
# Sum/avg over horizon steps
logprobs = logprobs.mean(-1) # -> (n_steps x n_envs)
@@ -460,6 +443,6 @@ class VPGDiffusion(DiffusionModel):
loss_actor = torch.mean(-logprobs * advantage)
# Train critic to predict state value
pred = self.critic(obs).squeeze()
pred = self.critic(cond).squeeze()
loss_critic = F.mse_loss(pred, reward)
return loss_actor, loss_critic, eta
+24 -22
View File
@@ -28,14 +28,11 @@ class EtaFixed(torch.nn.Module):
torch.tensor([2 * (base_eta - min_eta) / (max_eta - min_eta) - 1])
)
def __call__(self, x):
def __call__(self, cond):
"""Match input batch size, but do not depend on input"""
if isinstance(x, dict):
B = x["state"].shape[0]
device = x["state"].device
else:
B = x.size(0)
device = x.device
sample_data = cond["state"] if "state" in cond else cond["rgb"]
B = len(sample_data)
device = sample_data.device
eta_normalized = torch.tanh(self.eta_logit)
# map to min and max from [-1, 1]
@@ -64,14 +61,11 @@ class EtaAction(torch.nn.Module):
self.min = min_eta
self.max = max_eta
def __call__(self, x):
def __call__(self, cond):
"""Match input batch size, but do not depend on input"""
if isinstance(x, dict):
B = x["state"].shape[0]
device = x["state"].device
else:
B = x.size(0)
device = x.device
sample_data = cond["state"] if "state" in cond else cond["rgb"]
B = len(sample_data)
device = sample_data.device
eta_normalized = torch.tanh(self.eta_logit)
# map to min and max from [-1, 1]
@@ -109,13 +103,17 @@ class EtaState(torch.nn.Module):
torch.nn.init.xavier_normal_(m.weight, gain=gain)
m.bias.data.fill_(0)
def __call__(self, x):
if isinstance(x, dict):
def __call__(self, cond):
if "rgb" in cond:
raise NotImplementedError(
"State-based eta not implemented for image-based training!"
)
x = x.view(x.size(0), -1)
eta_res = self.mlp_res(x)
# flatten history
B = len(cond["state"])
state = cond["state"].view(B, -1)
# forward pass
eta_res = self.mlp_res(state)
eta_res = torch.tanh(eta_res) # [-1, 1]
eta = eta_res + self.base # [0, 2]
return torch.clamp(eta, self.min_res + self.base, self.max_res + self.base)
@@ -152,13 +150,17 @@ class EtaStateAction(torch.nn.Module):
torch.nn.init.xavier_normal_(m.weight, gain=gain)
m.bias.data.fill_(0)
def __call__(self, x):
if isinstance(x, dict):
def __call__(self, cond):
if "rgb" in cond:
raise NotImplementedError(
"State-action-based eta not implemented for image-based training!"
)
x = x.view(x.size(0), -1)
eta_res = self.mlp_res(x)
# flatten history
B = len(cond["state"])
state = cond["state"].view(B, -1)
# forward pass
eta_res = self.mlp_res(state)
eta_res = torch.tanh(eta_res) # [-1, 1]
eta = eta_res + self.base
return torch.clamp(eta, self.min_res + self.base, self.max_res + self.base)
+8 -5
View File
@@ -93,11 +93,12 @@ def get_likelihood_fn(
"""Compute an unbiased estimate to the log-likelihood in bits/dim.
Args:
model: A score model.
data: A PyTorch tensor. B x horizon x transition_dim
cond: dict with key state/rgb; more recent obs at the end
state: (B, To, Do)
data: (B x Ta x Da)
Returns:
logprob: B
logprob: (B,)
"""
shape = data.shape
B, H, A = shape
@@ -118,7 +119,9 @@ def get_likelihood_fn(
raise NotImplementedError(f"Hutchinson type {hutchinson_type} unknown.")
# repeat for expectation
cond_eps = cond.repeat_interleave(num_epsilon, dim=0)
cond_eps = {
key: cond[key].repeat_interleave(num_epsilon, dim=0) for key in cond
}
def ode_func(t, x):
x = x[:, :-1]
@@ -132,7 +135,7 @@ def get_likelihood_fn(
model_fn = model_ft
else:
model_fn = model
x = x.view(shape) # B x horizon x transition_dim
x = x.view(shape) # B x horizon x action_dim
drift = drift_fn(
model_fn,
x,
+46 -30
View File
@@ -25,6 +25,7 @@ class VisionDiffusionMLP(nn.Module):
transition_dim,
horizon_steps,
cond_dim,
img_cond_steps=1,
time_dim=16,
mlp_dims=[256, 256],
activation_type="Mish",
@@ -46,6 +47,8 @@ class VisionDiffusionMLP(nn.Module):
if augment:
self.aug = RandomShiftsAug(pad=4)
self.augment = augment
self.num_img = num_img
self.img_cond_steps = img_cond_steps
if spatial_emb > 0:
assert spatial_emb > 1, "this is the dimension"
if num_img > 1:
@@ -101,36 +104,44 @@ class VisionDiffusionMLP(nn.Module):
self,
x,
time,
cond=None,
cond: dict,
**kwargs,
):
"""
x: (B,T,obs_dim)
x: (B, Ta, Da)
time: (B,) or int, diffusion step
cond: dict (B,cond_step,cond_dim)
output: (B,T,input_dim)
cond: dict with key state/rgb; more recent obs at the end
state: (B, To, Do)
rgb: (B, To, C, H, W)
TODO long term: more flexible handling of cond
"""
# flatten T and input_dim
B, T, input_dim = x.shape
B, Ta, Da = x.shape
_, T_rgb, C, H, W = cond["rgb"].shape
# flatten chunk
x = x.view(B, -1)
# flatten cond_dim if exists
if cond["rgb"].ndim == 5:
rgb = einops.rearrange(cond["rgb"], "b d c h w -> (b d) c h w")
# flatten history
state = cond["state"].view(B, -1)
# Take recent images --- sometimes we want to use fewer img_cond_steps than cond_steps (e.g., 1 image but 3 prio)
rgb = cond["rgb"][:, -self.img_cond_steps :]
# concatenate images in cond by channels
if self.num_img > 1:
rgb = rgb.reshape(B, T_rgb, self.num_img, 3, H, W)
rgb = einops.rearrange(rgb, "b t n c h w -> b n (t c) h w")
else:
rgb = cond["rgb"]
if cond["state"].ndim == 3:
state = einops.rearrange(cond["state"], "b d c -> (b d) c")
else:
state = cond["state"]
rgb = einops.rearrange(rgb, "b t c h w -> b (t c) h w")
# convert rgb to float32 for augmentation
rgb = rgb.float()
# get vit output - pass in two images separately
if rgb.shape[1] == 6: # TODO: properly handle multiple images
rgb1 = rgb[:, :3]
rgb2 = rgb[:, 3:]
if self.num_img > 1: # TODO: properly handle multiple images
rgb1 = rgb[:, 0]
rgb2 = rgb[:, 1]
if self.augment:
rgb1 = self.aug(rgb1)
rgb2 = self.aug(rgb2)
@@ -141,7 +152,7 @@ class VisionDiffusionMLP(nn.Module):
feat = torch.cat([feat1, feat2], dim=-1)
else: # single image
if self.augment:
rgb = self.aug(rgb) # uint8 -> float32
rgb = self.aug(rgb)
feat = self.backbone(rgb)
# compress
@@ -159,7 +170,7 @@ class VisionDiffusionMLP(nn.Module):
# mlp
out = self.mlp_mean(x)
return out.view(B, T, input_dim)
return out.view(B, Ta, Da)
class DiffusionMLP(nn.Module):
@@ -210,27 +221,32 @@ class DiffusionMLP(nn.Module):
self,
x,
time,
cond=None,
cond,
**kwargs,
):
"""
x: (B,T,obs_dim)
x: (B, Ta, Da)
time: (B,) or int, diffusion step
cond: (B,cond_step,cond_dim)
output: (B,T,input_dim)
cond: dict with key state/rgb; more recent obs at the end
state: (B, To, Do)
"""
# flatten T and input_dim
B, T, input_dim = x.shape
B, Ta, Da = x.shape
# flatten chunk
x = x.view(B, -1)
cond = cond.view(B, -1) if cond is not None else None
# flatten history
state = cond["state"].view(B, -1)
# obs encoder
if hasattr(self, "cond_mlp"):
cond = self.cond_mlp(cond)
state = self.cond_mlp(state)
# append time and cond
time = time.view(B, 1)
time_emb = self.time_embedding(time).view(B, self.time_dim)
x = torch.cat([x, time_emb, cond], dim=-1)
x = torch.cat([x, time_emb, state], dim=-1)
# mlp
# mlp head
out = self.mlp_mean(x)
return out.view(B, T, input_dim)
return out.view(B, Ta, Da)
-6
View File
@@ -26,12 +26,6 @@ def extract(a, t, x_shape):
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
def apply_obs_conditioning(x, conditions, action_dim):
for t, val in conditions.items():
x[:, t, action_dim:] = val.clone()
return x
def make_timesteps(batch_size, i, device):
t = torch.full((batch_size,), i, device=device, dtype=torch.long)
return t
+13 -6
View File
@@ -270,15 +270,22 @@ class Unet1D(nn.Module):
**kwargs,
):
"""
x: (B,T,input_dim)
x: (B, Ta, act_dim)
time: (B,) or int, diffusion step
cond: (B,obs_step,cond_dim)
output: (B,T,input_dim)
cond: dict with key state/rgb; more recent obs at the end
state: (B, To, obs_dim)
"""
B = len(x)
# move chunk dim to the end
x = einops.rearrange(x, "b h t -> b t h")
cond = cond.view(cond.shape[0], -1)
# flatten history
state = cond["state"].view(B, -1)
# obs encoder
if hasattr(self, "cond_mlp"):
cond = self.cond_mlp(cond)
state = self.cond_mlp(state)
# 1. time
if not torch.is_tensor(time):
@@ -288,7 +295,7 @@ class Unet1D(nn.Module):
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
time = time.expand(x.shape[0])
global_feature = self.time_mlp(time)
global_feature = torch.cat([global_feature, cond], axis=-1)
global_feature = torch.cat([global_feature, state], axis=-1)
# encode local features
h_local = list()