squash commits
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user