release
This commit is contained in:
@@ -0,0 +1,318 @@
|
||||
"""
|
||||
Gaussian diffusion with DDPM and optionally DDIM sampling.
|
||||
|
||||
References:
|
||||
Diffuser: https://github.com/jannerm/diffuser
|
||||
Diffusion Policy: https://github.com/columbia-ai-robotics/diffusion_policy/blob/main/diffusion_policy/policy/diffusion_unet_lowdim_policy.py
|
||||
Annotated DDIM/DDPM: https://nn.labml.ai/diffusion/stable_diffusion/sampler/ddpm.html
|
||||
|
||||
"""
|
||||
|
||||
from typing import Optional, Union
|
||||
import logging
|
||||
import torch
|
||||
from torch import nn
|
||||
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")
|
||||
|
||||
|
||||
class DiffusionModel(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
network,
|
||||
horizon_steps,
|
||||
obs_dim,
|
||||
action_dim,
|
||||
transition_dim,
|
||||
network_path=None,
|
||||
cond_steps=1,
|
||||
device="cuda:0",
|
||||
# DDPM parameters
|
||||
denoising_steps=100,
|
||||
predict_epsilon=True,
|
||||
denoised_clip_value=1.0,
|
||||
# DDIM sampling
|
||||
use_ddim=False,
|
||||
ddim_discretize='uniform',
|
||||
ddim_steps=None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.device = device
|
||||
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
|
||||
|
||||
# Set up models
|
||||
self.network = network.to(device)
|
||||
if network_path is not None:
|
||||
checkpoint = torch.load(network_path, map_location=device, weights_only=True)
|
||||
if "ema" in checkpoint:
|
||||
self.load_state_dict(checkpoint["ema"], strict=False)
|
||||
logging.info("Loaded SL-trained policy from %s", network_path)
|
||||
else:
|
||||
self.load_state_dict(checkpoint["model"], strict=False)
|
||||
logging.info("Loaded RL-trained policy from %s", network_path)
|
||||
logging.info(
|
||||
f"Number of network parameters: {sum(p.numel() for p in self.parameters())}"
|
||||
)
|
||||
|
||||
"""
|
||||
DDPM parameters
|
||||
|
||||
"""
|
||||
"""
|
||||
βₜ
|
||||
"""
|
||||
self.betas = cosine_beta_schedule(denoising_steps).to(device)
|
||||
"""
|
||||
αₜ = 1 - βₜ
|
||||
"""
|
||||
self.alphas = 1.0 - self.betas
|
||||
"""
|
||||
α̅ₜ= ∏ᵗₛ₌₁ αₛ
|
||||
"""
|
||||
self.alphas_cumprod = torch.cumprod(self.alphas, axis=0)
|
||||
"""
|
||||
α̅ₜ₋₁
|
||||
"""
|
||||
self.alphas_cumprod_prev = torch.cat([torch.ones(1).to(device), self.alphas_cumprod[:-1]])
|
||||
"""
|
||||
√ α̅ₜ
|
||||
"""
|
||||
self.sqrt_alphas_cumprod = torch.sqrt(self.alphas_cumprod)
|
||||
"""
|
||||
√ 1-α̅ₜ
|
||||
"""
|
||||
self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - self.alphas_cumprod)
|
||||
"""
|
||||
√ 1\α̅ₜ
|
||||
"""
|
||||
self.sqrt_recip_alphas_cumprod = torch.sqrt(1.0 / self.alphas_cumprod)
|
||||
"""
|
||||
√ 1\α̅ₜ-1
|
||||
"""
|
||||
self.sqrt_recipm1_alphas_cumprod = torch.sqrt(1.0 / self.alphas_cumprod - 1)
|
||||
"""
|
||||
β̃ₜ = σₜ² = βₜ (1-α̅ₜ₋₁)/(1-α̅ₜ)
|
||||
"""
|
||||
self.ddpm_var = (
|
||||
self.betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
|
||||
)
|
||||
self.ddpm_logvar_clipped = torch.log(torch.clamp(self.ddpm_var, min=1e-20))
|
||||
"""
|
||||
μₜ = β̃ₜ √ α̅ₜ₋₁/(1-α̅ₜ)x₀ + √ αₜ (1-α̅ₜ₋₁)/(1-α̅ₜ)xₜ
|
||||
"""
|
||||
self.ddpm_mu_coef1 = self.betas * torch.sqrt(self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
|
||||
self.ddpm_mu_coef2 = (1.0 - self.alphas_cumprod_prev) * torch.sqrt(self.alphas) / (1.0 - self.alphas_cumprod)
|
||||
|
||||
"""
|
||||
DDIM parameters
|
||||
|
||||
In DDIM paper https://arxiv.org/pdf/2010.02502, alpha is alpha_cumprod in DDPM https://arxiv.org/pdf/2102.09672
|
||||
"""
|
||||
if use_ddim:
|
||||
assert predict_epsilon, "DDIM requires predicting epsilon for now."
|
||||
if ddim_discretize == 'uniform': # use the HF "leading" style
|
||||
step_ratio = self.denoising_steps // ddim_steps
|
||||
self.ddim_t = torch.arange(0, ddim_steps, device=self.device) * step_ratio
|
||||
else:
|
||||
raise 'Unknown discretization method for DDIM.'
|
||||
self.ddim_alphas = self.alphas_cumprod[self.ddim_t].clone().to(torch.float32)
|
||||
self.ddim_alphas_sqrt = torch.sqrt(self.ddim_alphas)
|
||||
self.ddim_alphas_prev = torch.cat([
|
||||
torch.tensor([1.]).to(torch.float32).to(self.device),
|
||||
self.alphas_cumprod[self.ddim_t[:-1]]])
|
||||
self.ddim_sqrt_one_minus_alphas = (1. - self.ddim_alphas) ** .5
|
||||
|
||||
# Initialize fixed sigmas for inference - eta=0
|
||||
ddim_eta = 0
|
||||
self.ddim_sigmas = (ddim_eta * \
|
||||
((1 - self.ddim_alphas_prev) / (1 - self.ddim_alphas) * \
|
||||
(1 - self.ddim_alphas / self.ddim_alphas_prev)) ** .5)
|
||||
|
||||
# Flip all
|
||||
self.ddim_t = torch.flip(self.ddim_t, [0])
|
||||
self.ddim_alphas = torch.flip(self.ddim_alphas, [0])
|
||||
self.ddim_alphas_sqrt = torch.flip(self.ddim_alphas_sqrt, [0])
|
||||
self.ddim_alphas_prev = torch.flip(self.ddim_alphas_prev, [0])
|
||||
self.ddim_sqrt_one_minus_alphas = torch.flip(self.ddim_sqrt_one_minus_alphas, [0])
|
||||
self.ddim_sigmas = torch.flip(self.ddim_sigmas, [0])
|
||||
|
||||
# ---------- Sampling ----------#
|
||||
|
||||
def p_mean_var(self, x, t, cond=None, index=None):
|
||||
noise = self.network(x, t, cond=cond)
|
||||
|
||||
# Predict x_0
|
||||
if self.predict_epsilon:
|
||||
if self.use_ddim:
|
||||
"""
|
||||
x₀ = (xₜ - √ (1-αₜ) ε )/ √ αₜ
|
||||
"""
|
||||
alpha = extract(self.ddim_alphas, index, x.shape)
|
||||
alpha_prev = extract(self.ddim_alphas_prev, index, x.shape)
|
||||
sqrt_one_minus_alpha = extract(self.ddim_sqrt_one_minus_alphas, index, x.shape)
|
||||
x_recon = (x - sqrt_one_minus_alpha * noise) / (alpha ** 0.5)
|
||||
else:
|
||||
"""
|
||||
x₀ = √ 1\α̅ₜ xₜ - √ 1\α̅ₜ-1 ε
|
||||
"""
|
||||
x_recon = (
|
||||
extract(self.sqrt_recip_alphas_cumprod, t, x.shape) * x
|
||||
- extract(self.sqrt_recipm1_alphas_cumprod, t, x.shape) * noise
|
||||
)
|
||||
else: # directly predicting x₀
|
||||
x_recon = noise
|
||||
if self.denoised_clip_value is not None:
|
||||
x_recon.clamp_(-self.denoised_clip_value, self.denoised_clip_value)
|
||||
if self.use_ddim:
|
||||
# re-calculate noise based on clamped x_recon - default to false in HF, but let's use it here
|
||||
noise = (x - alpha ** (0.5) * x_recon) / sqrt_one_minus_alpha
|
||||
|
||||
# Get mu
|
||||
if self.use_ddim:
|
||||
"""
|
||||
μ = √ αₜ₋₁ x₀ + √(1-αₜ₋₁ - σₜ²) ε
|
||||
|
||||
var should be zero here as self.ddim_eta=0
|
||||
"""
|
||||
sigma = extract(self.ddim_sigmas, index, x.shape)
|
||||
dir_xt = (1. - alpha_prev - sigma ** 2).sqrt() * noise
|
||||
mu = (alpha_prev ** 0.5) * x_recon + dir_xt
|
||||
var = sigma ** 2
|
||||
logvar = torch.log(var)
|
||||
else:
|
||||
"""
|
||||
μₜ = β̃ₜ √ α̅ₜ₋₁/(1-α̅ₜ)x₀ + √ αₜ (1-α̅ₜ₋₁)/(1-α̅ₜ)xₜ
|
||||
"""
|
||||
mu = (
|
||||
extract(self.ddpm_mu_coef1, t, x.shape) * x_recon
|
||||
+ extract(self.ddpm_mu_coef2, t, x.shape) * x
|
||||
)
|
||||
logvar = extract(
|
||||
self.ddpm_logvar_clipped, t, x.shape
|
||||
)
|
||||
return mu, logvar
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
cond: Optional[torch.Tensor],
|
||||
return_chain=True,
|
||||
):
|
||||
"""
|
||||
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)
|
||||
|
||||
# ---------- Supervised training ----------#
|
||||
|
||||
def loss(self, x, *args):
|
||||
batch_size = len(x)
|
||||
t = torch.randint(
|
||||
0, self.denoising_steps, (batch_size,), device=x.device
|
||||
).long()
|
||||
return self.p_losses(x, *args, t)
|
||||
|
||||
def p_losses(
|
||||
self,
|
||||
x_start,
|
||||
obs_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
|
||||
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)
|
||||
x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise)
|
||||
|
||||
# Predict
|
||||
x_recon = self.network(x_noisy, t, cond=cond)
|
||||
if self.predict_epsilon:
|
||||
return F.mse_loss(x_recon, noise, reduction="mean")
|
||||
else:
|
||||
return F.mse_loss(x_recon, x_noisy, reduction="mean")
|
||||
|
||||
def q_sample(self, x_start, t, noise=None):
|
||||
"""
|
||||
q(xₜ | x₀) = 𝒩(xₜ; √ α̅ₜ x₀, (1-α̅ₜ)I)
|
||||
xₜ = √ α̅ₜ xₒ + √ (1-α̅ₜ) ε
|
||||
"""
|
||||
if noise is None:
|
||||
device = x_start.device
|
||||
noise = torch.randn_like(x_start, device=device)
|
||||
return (
|
||||
extract(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start
|
||||
+ extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise
|
||||
)
|
||||
@@ -0,0 +1,34 @@
|
||||
"""
|
||||
Advantage-weighted regression (AWR) for diffusion policy.
|
||||
|
||||
"""
|
||||
|
||||
import logging
|
||||
import torch
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
from model.diffusion.diffusion_rwr import RWRDiffusion
|
||||
|
||||
|
||||
class AWRDiffusion(RWRDiffusion):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
actor,
|
||||
critic,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(network=actor, **kwargs)
|
||||
self.critic = critic.to(self.device)
|
||||
|
||||
# assign actor
|
||||
self.actor = self.network
|
||||
|
||||
def loss_critic(self, obs, advantages):
|
||||
# get advantage
|
||||
adv = self.critic(obs)
|
||||
|
||||
# Update critic
|
||||
loss_critic = torch.mean((adv - advantages) ** 2)
|
||||
return loss_critic
|
||||
@@ -0,0 +1,118 @@
|
||||
"""
|
||||
Actor and Critic models for model-free online RL with DIffusion POlicy (DIPO).
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
import logging
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
from model.diffusion.diffusion import DiffusionModel
|
||||
from model.diffusion.sampling import make_timesteps
|
||||
|
||||
|
||||
class DIPODiffusion(DiffusionModel):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
actor,
|
||||
critic,
|
||||
use_ddim=False,
|
||||
randn_clip_value=10,
|
||||
clamp_action=False,
|
||||
min_sampling_denoising_std=0.1,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(network=actor, use_ddim=use_ddim, **kwargs)
|
||||
assert not self.use_ddim, "DQL does not support DDIM"
|
||||
self.critic = critic.to(self.device)
|
||||
|
||||
# reassign actor
|
||||
self.actor = self.network
|
||||
|
||||
# Minimum std used in denoising process when sampling action - helps exploration
|
||||
self.min_sampling_denoising_std = min_sampling_denoising_std
|
||||
|
||||
# For each denoising step, we clip sampled randn (from standard deviation) such that the sampled action is not too far away from mean
|
||||
self.randn_clip_value = randn_clip_value
|
||||
|
||||
# Whether to clamp sampled action between [-1, 1]
|
||||
self.clamp_action = clamp_action
|
||||
|
||||
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)
|
||||
|
||||
# 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)
|
||||
next_q = torch.min(next_q1, next_q2)
|
||||
|
||||
# terminal state mask
|
||||
mask = 1 - dones
|
||||
|
||||
# flatten
|
||||
rewards = rewards.view(-1)
|
||||
next_q = next_q.view(-1)
|
||||
mask = mask.view(-1)
|
||||
|
||||
# target value
|
||||
target_q = rewards + gamma * next_q * mask
|
||||
|
||||
# Update critic
|
||||
loss_critic = torch.mean((current_q1 - target_q) ** 2) + torch.mean(
|
||||
(current_q2 - target_q) ** 2
|
||||
)
|
||||
|
||||
return loss_critic
|
||||
|
||||
# override
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
cond,
|
||||
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]
|
||||
|
||||
# Loop
|
||||
x = torch.randn((B, self.horizon_steps, self.transition_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)
|
||||
mean, logvar = self.p_mean_var(
|
||||
x=x,
|
||||
t=t_b,
|
||||
cond=cond,
|
||||
)
|
||||
std = torch.exp(0.5 * logvar)
|
||||
|
||||
# Determine the noise level
|
||||
if deterministic and t == 0:
|
||||
std = torch.zeros_like(std)
|
||||
elif deterministic: # For DDPM, sample with noise
|
||||
std = torch.clip(std, min=1e-3)
|
||||
else:
|
||||
std = torch.clip(std, min=self.min_sampling_denoising_std)
|
||||
noise = torch.randn_like(x).clamp_(
|
||||
-self.randn_clip_value, self.randn_clip_value
|
||||
)
|
||||
x = mean + std * noise
|
||||
|
||||
# clamp action at final step
|
||||
if self.clamp_action and i == len(t_all) - 1:
|
||||
x = torch.clamp(x, -1, 1)
|
||||
return x
|
||||
@@ -0,0 +1,173 @@
|
||||
"""
|
||||
Diffusion Q-Learning (DQL)
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
import logging
|
||||
import numpy as np
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
from model.diffusion.diffusion import DiffusionModel
|
||||
from model.diffusion.sampling import make_timesteps
|
||||
|
||||
|
||||
class DQLDiffusion(DiffusionModel):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
actor,
|
||||
critic,
|
||||
use_ddim=False,
|
||||
randn_clip_value=10,
|
||||
clamp_action=False,
|
||||
min_sampling_denoising_std=0.1,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(network=actor, use_ddim=use_ddim, **kwargs)
|
||||
assert not self.use_ddim, "DQL does not support DDIM"
|
||||
self.critic = critic.to(self.device)
|
||||
|
||||
# reassign actor
|
||||
self.actor = self.network
|
||||
|
||||
# Minimum std used in denoising process when sampling action - helps exploration
|
||||
self.min_sampling_denoising_std = min_sampling_denoising_std
|
||||
|
||||
# For each denoising step, we clip sampled randn (from standard deviation) such that the sampled action is not too far away from mean
|
||||
self.randn_clip_value = randn_clip_value
|
||||
|
||||
# Whether to clamp sampled action between [-1, 1]
|
||||
self.clamp_action = clamp_action
|
||||
|
||||
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)
|
||||
|
||||
# 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)
|
||||
next_q = torch.min(next_q1, next_q2)
|
||||
|
||||
# terminal state mask
|
||||
mask = 1 - dones
|
||||
|
||||
# flatten
|
||||
rewards = rewards.view(-1)
|
||||
next_q = next_q.view(-1)
|
||||
mask = mask.view(-1)
|
||||
|
||||
# target value
|
||||
target_q = rewards + gamma * next_q * mask
|
||||
|
||||
# Update critic
|
||||
loss_critic = torch.mean((current_q1 - target_q) ** 2) + torch.mean(
|
||||
(current_q2 - target_q) ** 2
|
||||
)
|
||||
|
||||
return loss_critic
|
||||
|
||||
def loss_actor(self, obs, actions, q1, q2, eta):
|
||||
bc_loss = self.loss(actions, {0: obs})
|
||||
if np.random.uniform() > 0.5:
|
||||
q_loss = -q1.mean() / q2.abs().mean().detach()
|
||||
else:
|
||||
q_loss = -q2.mean() / q1.abs().mean().detach()
|
||||
actor_loss = bc_loss + eta * q_loss
|
||||
return actor_loss
|
||||
|
||||
# override
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
cond,
|
||||
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]
|
||||
|
||||
# Loop
|
||||
x = torch.randn((B, self.horizon_steps, self.transition_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)
|
||||
mean, logvar = self.p_mean_var(
|
||||
x=x,
|
||||
t=t_b,
|
||||
cond=cond,
|
||||
)
|
||||
std = torch.exp(0.5 * logvar)
|
||||
|
||||
# Determine the noise level
|
||||
if deterministic and t == 0:
|
||||
std = torch.zeros_like(std)
|
||||
elif deterministic: # For DDPM, sample with noise
|
||||
std = torch.clip(std, min=1e-3)
|
||||
else:
|
||||
std = torch.clip(std, min=self.min_sampling_denoising_std)
|
||||
noise = torch.randn_like(x).clamp_(
|
||||
-self.randn_clip_value, self.randn_clip_value
|
||||
)
|
||||
x = mean + std * noise
|
||||
|
||||
# clamp action at final step
|
||||
if self.clamp_action and i == len(t_all) - 1:
|
||||
x = torch.clamp(x, -1, 1)
|
||||
return x
|
||||
|
||||
def forward_train(
|
||||
self,
|
||||
cond,
|
||||
deterministic=False,
|
||||
):
|
||||
"""
|
||||
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]
|
||||
|
||||
# Loop
|
||||
x = torch.randn((B, self.horizon_steps, self.transition_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)
|
||||
mean, logvar = self.p_mean_var(
|
||||
x=x,
|
||||
t=t_b,
|
||||
cond=cond,
|
||||
)
|
||||
std = torch.exp(0.5 * logvar)
|
||||
|
||||
# Determine the noise level
|
||||
if deterministic and t == 0:
|
||||
std = torch.zeros_like(std)
|
||||
elif deterministic: # For DDPM, sample with noise
|
||||
std = torch.clip(std, min=1e-3)
|
||||
else:
|
||||
std = torch.clip(std, min=self.min_sampling_denoising_std)
|
||||
noise = torch.randn_like(x).clamp_(
|
||||
-self.randn_clip_value, self.randn_clip_value
|
||||
)
|
||||
x = mean + std * noise
|
||||
|
||||
# clamp action at final step
|
||||
if self.clamp_action and i == len(t_all) - 1:
|
||||
x = torch.clamp(x, -1, 1)
|
||||
return x
|
||||
@@ -0,0 +1,197 @@
|
||||
"""
|
||||
Implicit diffusion Q-learning (IDQL) for diffusion policy.
|
||||
|
||||
"""
|
||||
|
||||
import logging
|
||||
import torch
|
||||
import einops
|
||||
import copy
|
||||
|
||||
import torch.nn.functional as F
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
from model.diffusion.diffusion_rwr import RWRDiffusion
|
||||
|
||||
|
||||
def expectile_loss(diff, expectile=0.8):
|
||||
weight = torch.where(diff > 0, expectile, (1 - expectile))
|
||||
return weight * (diff**2)
|
||||
|
||||
|
||||
def soft_update(target, source, tau):
|
||||
for target_param, param in zip(target.parameters(), source.parameters()):
|
||||
target_param.data.copy_(target_param.data * (1.0 - tau) + param.data * tau)
|
||||
|
||||
|
||||
class IDQLDiffusion(RWRDiffusion):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
actor,
|
||||
critic_q,
|
||||
critic_v,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(network=actor, **kwargs)
|
||||
self.critic_q = critic_q.to(self.device)
|
||||
self.target_q = copy.deepcopy(critic_q)
|
||||
self.critic_v = critic_v.to(self.device)
|
||||
|
||||
# assign actor
|
||||
self.actor = self.network
|
||||
|
||||
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)
|
||||
q = torch.min(current_q1, current_q2)
|
||||
|
||||
# get the current V-function
|
||||
v = self.critic_v(obs).reshape(-1)
|
||||
|
||||
# compute advantage
|
||||
adv = q - v
|
||||
|
||||
return adv
|
||||
|
||||
def loss_critic_v(self, obs, actions):
|
||||
|
||||
adv = self.compute_advantages(obs, actions)
|
||||
|
||||
# get the value loss
|
||||
v_loss = expectile_loss(adv).mean()
|
||||
|
||||
return v_loss
|
||||
|
||||
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)
|
||||
|
||||
# get the next V-function
|
||||
with torch.no_grad(): # no gradients for value function when we update q function
|
||||
next_v = self.critic_v(next_obs)
|
||||
|
||||
# terminal state mask
|
||||
mask = 1 - dones
|
||||
|
||||
# flatten
|
||||
rewards = rewards.view(-1)
|
||||
next_v = next_v.view(-1)
|
||||
mask = mask.view(-1)
|
||||
|
||||
# target value
|
||||
discounted_q = rewards + gamma * next_v * mask
|
||||
|
||||
# Update critic
|
||||
q_loss = torch.mean((current_q1 - discounted_q) ** 2) + torch.mean(
|
||||
(current_q2 - discounted_q) ** 2
|
||||
)
|
||||
|
||||
return q_loss
|
||||
|
||||
def update_target_critic(self, tau):
|
||||
soft_update(self.target_q, self.critic_q, tau)
|
||||
|
||||
@torch.no_grad()
|
||||
def forward( # override
|
||||
self,
|
||||
cond,
|
||||
deterministic=False,
|
||||
num_sample=10,
|
||||
critic_hyperparam=0.7, # sampling weight for implicit policy
|
||||
use_expectile_exploration=True,
|
||||
):
|
||||
# repeat obs num_sample times along dim 0
|
||||
cond_shape_repeat_dims = tuple(1 for _ in cond.shape)
|
||||
B, T, D = cond.shape
|
||||
S = num_sample
|
||||
cond_repeat = cond[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,
|
||||
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)
|
||||
q = torch.min(current_q1, current_q2)
|
||||
q = q.view(S, B)
|
||||
|
||||
# Use argmax
|
||||
if deterministic or (not use_expectile_exploration):
|
||||
# gather the best sample -- filter out suboptimal Q during inference
|
||||
best_indices = q.argmax(0)
|
||||
samples_expanded = samples.view(S, B, H, A)
|
||||
|
||||
# dummy dimension @ dim 0 for batched indexing
|
||||
sample_indices = best_indices[None, :, None, None] # [1, B, 1, 1]
|
||||
sample_indices = sample_indices.repeat(S, 1, H, A)
|
||||
|
||||
samples_best = torch.gather(samples_expanded, 0, sample_indices)
|
||||
# Sample as an implicit policy for exploration
|
||||
else:
|
||||
# get the current value function for probabilistic exploration
|
||||
current_v = self.critic_v(cond_repeat)
|
||||
v = current_v.view(S, B)
|
||||
adv = q - v
|
||||
|
||||
# Compute weights for sampling
|
||||
samples_expanded = samples.view(S, B, H, A)
|
||||
|
||||
# expectile exploration policy
|
||||
tau_weights = torch.where(adv > 0, critic_hyperparam, 1 - critic_hyperparam)
|
||||
tau_weights = tau_weights / tau_weights.sum(0) # normalize
|
||||
|
||||
# select a sample from DP probabilistically -- sample index per batch and compile
|
||||
sample_indices = torch.multinomial(tau_weights.T, 1) # [B, 1]
|
||||
|
||||
# dummy dimension @ dim 0 for batched indexing
|
||||
sample_indices = sample_indices[None, :, None] # [1, B, 1, 1]
|
||||
sample_indices = sample_indices.repeat(S, 1, H, A)
|
||||
|
||||
samples_best = torch.gather(samples_expanded, 0, sample_indices)
|
||||
|
||||
# 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()
|
||||
@@ -0,0 +1,193 @@
|
||||
"""
|
||||
DPPO: Diffusion Policy Policy Optimization.
|
||||
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
import torch
|
||||
import logging
|
||||
import math
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from model.diffusion.diffusion_vpg import VPGDiffusion
|
||||
|
||||
|
||||
class PPODiffusion(VPGDiffusion):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
gamma_denoising: float,
|
||||
clip_ploss_coef: float,
|
||||
clip_ploss_coef_base: float = 1e-3,
|
||||
clip_ploss_coef_rate: float = 3,
|
||||
clip_vloss_coef: Optional[float] = None,
|
||||
clip_advantage_lower_quantile: float = 0,
|
||||
clip_advantage_upper_quantile: float = 1,
|
||||
norm_adv: bool = True,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
# Whether to normalize advantages within batch
|
||||
self.norm_adv = norm_adv
|
||||
|
||||
# Clipping value for policy loss
|
||||
self.clip_ploss_coef = clip_ploss_coef
|
||||
self.clip_ploss_coef_base = clip_ploss_coef_base
|
||||
self.clip_ploss_coef_rate = clip_ploss_coef_rate
|
||||
|
||||
# Clipping value for value loss
|
||||
self.clip_vloss_coef = clip_vloss_coef
|
||||
|
||||
# Discount factor for diffusion MDP
|
||||
self.gamma_denoising = gamma_denoising
|
||||
|
||||
# Quantiles for clipping advantages
|
||||
self.clip_advantage_lower_quantile = clip_advantage_lower_quantile
|
||||
self.clip_advantage_upper_quantile = clip_advantage_upper_quantile
|
||||
|
||||
def loss(
|
||||
self,
|
||||
obs,
|
||||
chains,
|
||||
returns,
|
||||
oldvalues,
|
||||
advantages,
|
||||
oldlogprobs,
|
||||
use_bc_loss=False,
|
||||
reward_horizon=4,
|
||||
):
|
||||
"""
|
||||
PPO loss
|
||||
|
||||
obs: (B, obs_step, obs_dim)
|
||||
chains: (B, num_denoising_step+1, horizon_step, action_dim)
|
||||
returns: (B, )
|
||||
values: (B, )
|
||||
advantages: (B,)
|
||||
oldlogprobs: (B, num_denoising_step, horizon_step, action_dim)
|
||||
use_bc_loss: 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
|
||||
newlogprobs, eta = self.get_logprobs(
|
||||
obs,
|
||||
chains,
|
||||
get_ent=True,
|
||||
)
|
||||
entropy_loss = -eta.mean()
|
||||
newlogprobs = newlogprobs.clamp(min=-5, max=2)
|
||||
oldlogprobs = oldlogprobs.clamp(min=-5, max=2)
|
||||
|
||||
# only backpropagate through the earlier steps (e.g., ones actually executed in the environment)
|
||||
newlogprobs = newlogprobs[:, :reward_horizon, :]
|
||||
oldlogprobs = oldlogprobs[:, :, :reward_horizon, :]
|
||||
|
||||
# Get the logprobs - batch over B and denoising steps
|
||||
newlogprobs = newlogprobs.mean(dim=(-1, -2)).view(-1)
|
||||
oldlogprobs = oldlogprobs.mean(dim=(-1, -2)).view(-1)
|
||||
|
||||
bc_loss = 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,
|
||||
return_chain=True,
|
||||
use_base_policy=True,
|
||||
)
|
||||
# Get logprobs of teacher actions under this policy
|
||||
bc_logprobs = self.get_logprobs(
|
||||
obs,
|
||||
samples.chains, # n_env x denoising x horizon x act
|
||||
get_ent=False,
|
||||
use_base_policy=False,
|
||||
)
|
||||
bc_logprobs = bc_logprobs.clamp(min=-5, max=2)
|
||||
bc_logprobs = bc_logprobs.mean(dim=(-1, -2)).view(-1)
|
||||
bc_loss = -bc_logprobs.mean()
|
||||
|
||||
# normalize advantages
|
||||
if self.norm_adv:
|
||||
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
|
||||
|
||||
# Clip advantages by 5th and 95th percentile
|
||||
advantage_min = torch.quantile(advantages, self.clip_advantage_lower_quantile)
|
||||
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))]
|
||||
).to(self.device)
|
||||
discount = discount.repeat(len(advantages) // self.ft_denoising_steps)
|
||||
advantages *= discount
|
||||
|
||||
# get ratio
|
||||
logratio = newlogprobs - oldlogprobs
|
||||
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
|
||||
if self.ft_denoising_steps > 1:
|
||||
clip_ploss_coef = self.clip_ploss_coef_base + (
|
||||
self.clip_ploss_coef - self.clip_ploss_coef_base
|
||||
) * (torch.exp(self.clip_ploss_coef_rate * t) - 1) / (
|
||||
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
|
||||
)
|
||||
|
||||
# get kl difference and whether value clipped
|
||||
with torch.no_grad():
|
||||
# old_approx_kl: the approximate Kullback–Leibler divergence, measured by (-logratio).mean(), which corresponds to the k1 estimator in John Schulman’s blog post on approximating KL http://joschu.net/blog/kl-approx.html
|
||||
# approx_kl: better alternative to old_approx_kl measured by (logratio.exp() - 1) - logratio, which corresponds to the k3 estimator in approximating KL http://joschu.net/blog/kl-approx.html
|
||||
# old_approx_kl = (-logratio).mean()
|
||||
approx_kl = ((ratio - 1) - logratio).mean()
|
||||
clipfrac = ((ratio - 1.0).abs() > clip_ploss_coef).float().mean().item()
|
||||
|
||||
# Policy loss with clipping
|
||||
pg_loss1 = -advantages * ratio
|
||||
pg_loss2 = -advantages * torch.clamp(
|
||||
ratio, 1 - clip_ploss_coef, 1 + clip_ploss_coef
|
||||
)
|
||||
pg_loss = torch.max(pg_loss1, pg_loss2).mean()
|
||||
|
||||
# Value loss optionally with clipping
|
||||
newvalues = self.critic(obs).view(-1)
|
||||
if self.clip_vloss_coef is not None:
|
||||
v_loss_unclipped = (newvalues - returns) ** 2
|
||||
v_clipped = oldvalues + torch.clamp(
|
||||
newvalues - oldvalues,
|
||||
-self.clip_vloss_coef,
|
||||
self.clip_vloss_coef,
|
||||
)
|
||||
v_loss_clipped = (v_clipped - returns) ** 2
|
||||
v_loss_max = torch.max(v_loss_unclipped, v_loss_clipped)
|
||||
v_loss = 0.5 * v_loss_max.mean()
|
||||
else:
|
||||
v_loss = 0.5 * ((newvalues - returns) ** 2).mean()
|
||||
return (
|
||||
pg_loss,
|
||||
entropy_loss,
|
||||
v_loss,
|
||||
clipfrac,
|
||||
approx_kl.item(),
|
||||
ratio.mean().item(),
|
||||
bc_loss,
|
||||
eta.mean().item(),
|
||||
)
|
||||
@@ -0,0 +1,139 @@
|
||||
"""
|
||||
Diffusion policy gradient with exact likelihood estimation.
|
||||
|
||||
Based on score_sde_pytorch https://github.com/yang-song/score_sde_pytorch
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
import logging
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from .diffusion_ppo import PPODiffusion
|
||||
from .exact_likelihood import get_likelihood_fn
|
||||
|
||||
|
||||
class PPOExactDiffusion(PPODiffusion):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
sde,
|
||||
sde_hutchinson_type="Rademacher",
|
||||
sde_rtol=1e-4,
|
||||
sde_atol=1e-4,
|
||||
sde_eps=1e-4,
|
||||
sde_step_size=1e-3,
|
||||
sde_method="RK23",
|
||||
sde_continuous=False,
|
||||
sde_probability_flow=False,
|
||||
sde_num_epsilon=1,
|
||||
sde_min_beta=1e-2,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
self.sde = sde
|
||||
self.sde.set_betas(
|
||||
self.betas,
|
||||
sde_min_beta,
|
||||
)
|
||||
|
||||
# set up likelihood function
|
||||
self.likelihood_fn = get_likelihood_fn(
|
||||
sde,
|
||||
hutchinson_type=sde_hutchinson_type,
|
||||
rtol=sde_rtol,
|
||||
atol=sde_atol,
|
||||
eps=sde_eps,
|
||||
step_size=sde_step_size,
|
||||
method=sde_method,
|
||||
continuous=sde_continuous,
|
||||
probability_flow=sde_probability_flow,
|
||||
predict_epsilon=self.predict_epsilon,
|
||||
num_epsilon=sde_num_epsilon,
|
||||
)
|
||||
|
||||
def get_exact_logprobs(self, obs, samples):
|
||||
"""Use torchdiffeq
|
||||
|
||||
samples: B x horizon x transition_dim
|
||||
"""
|
||||
# TODO: image input
|
||||
cond = obs.reshape(-1, self.obs_dim)
|
||||
return self.likelihood_fn(
|
||||
self.actor,
|
||||
self.actor_ft,
|
||||
samples,
|
||||
self.denoising_steps,
|
||||
self.ft_denoising_steps,
|
||||
cond=cond,
|
||||
)
|
||||
|
||||
def loss(
|
||||
self,
|
||||
obs,
|
||||
samples,
|
||||
returns,
|
||||
oldvalues,
|
||||
advantages,
|
||||
oldlogprobs,
|
||||
use_bc_loss=False,
|
||||
**kwargs,
|
||||
):
|
||||
# Get new logprobs for final x
|
||||
newlogprobs = self.get_exact_logprobs(obs, samples)
|
||||
newlogprobs = newlogprobs.clamp(min=-5, max=2)
|
||||
oldlogprobs = oldlogprobs.clamp(min=-5, max=2)
|
||||
|
||||
bc_loss = 0
|
||||
if use_bc_loss:
|
||||
raise NotImplementedError
|
||||
|
||||
# get ratio
|
||||
logratio = newlogprobs - oldlogprobs
|
||||
ratio = logratio.exp()
|
||||
|
||||
# get kl difference and whether value clipped
|
||||
with torch.no_grad():
|
||||
# old_approx_kl: the approximate Kullback–Leibler divergence, measured by (-logratio).mean(), which corresponds to the k1 estimator in John Schulman’s blog post on approximating KL http://joschu.net/blog/kl-approx.html
|
||||
# approx_kl: better alternative to old_approx_kl measured by (logratio.exp() - 1) - logratio, which corresponds to the k3 estimator in approximating KL http://joschu.net/blog/kl-approx.html
|
||||
# old_approx_kl = (-logratio).mean()
|
||||
approx_kl = ((ratio - 1) - logratio).mean()
|
||||
clipfrac = (
|
||||
((ratio - 1.0).abs() > self.clip_ploss_coef).float().mean().item()
|
||||
)
|
||||
|
||||
# normalize advantages
|
||||
if self.norm_adv:
|
||||
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
|
||||
|
||||
# Policy loss with clipping
|
||||
pg_loss1 = -advantages * ratio
|
||||
pg_loss2 = -advantages * torch.clamp(
|
||||
ratio, 1 - self.clip_ploss_coef, 1 + self.clip_ploss_coef
|
||||
)
|
||||
pg_loss = torch.max(pg_loss1, pg_loss2).mean()
|
||||
|
||||
# Value loss optionally with clipping
|
||||
newvalues = self.critic(obs).view(-1)
|
||||
if self.clip_vloss_coef is not None:
|
||||
v_loss_unclipped = (newvalues - returns) ** 2
|
||||
v_clipped = oldvalues + torch.clamp(
|
||||
newvalues - oldvalues,
|
||||
-self.clip_vloss_coef,
|
||||
self.clip_vloss_coef,
|
||||
)
|
||||
v_loss_clipped = (v_clipped - returns) ** 2
|
||||
v_loss_max = torch.max(v_loss_unclipped, v_loss_clipped)
|
||||
v_loss = 0.5 * v_loss_max.mean()
|
||||
else:
|
||||
v_loss = 0.5 * ((newvalues - returns) ** 2).mean()
|
||||
|
||||
# entropy is maximized - only effective if residual is learned
|
||||
return (
|
||||
pg_loss,
|
||||
v_loss,
|
||||
clipfrac,
|
||||
approx_kl.item(),
|
||||
ratio.mean().item(),
|
||||
bc_loss,
|
||||
)
|
||||
@@ -0,0 +1,110 @@
|
||||
"""
|
||||
QSM (Q-Score Matching) for diffusion policy.
|
||||
|
||||
"""
|
||||
|
||||
import logging
|
||||
import torch
|
||||
import copy
|
||||
|
||||
import torch.nn.functional as F
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
from model.diffusion.diffusion_rwr import RWRDiffusion
|
||||
|
||||
|
||||
def expectile_loss(diff, expectile=0.8):
|
||||
weight = torch.where(diff > 0, expectile, (1 - expectile))
|
||||
return weight * (diff**2)
|
||||
|
||||
|
||||
def soft_update(target, source, tau):
|
||||
for target_param, param in zip(target.parameters(), source.parameters()):
|
||||
target_param.data.copy_(target_param.data * (1.0 - tau) + param.data * tau)
|
||||
|
||||
|
||||
class QSMDiffusion(RWRDiffusion):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
actor,
|
||||
critic,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(network=actor, **kwargs)
|
||||
self.critic_q = critic.to(self.device)
|
||||
self.target_q = copy.deepcopy(critic)
|
||||
|
||||
# assign actor
|
||||
self.actor = self.network
|
||||
|
||||
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)
|
||||
|
||||
# Forward process
|
||||
noise = torch.randn_like(x_start, device=device)
|
||||
t = torch.randint(
|
||||
0, self.denoising_steps, (B,), device=device
|
||||
).long() # sample random denoising time index
|
||||
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)
|
||||
|
||||
# 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_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)
|
||||
|
||||
# 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)
|
||||
|
||||
# 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)
|
||||
with torch.no_grad():
|
||||
next_q1, next_q2 = self.target_q(next_obs, next_actions_flat)
|
||||
next_q = torch.min(next_q1, next_q2)
|
||||
|
||||
# terminal state mask
|
||||
mask = 1 - dones
|
||||
|
||||
# flatten
|
||||
rewards = rewards.view(-1)
|
||||
next_q = next_q.view(-1)
|
||||
mask = mask.view(-1)
|
||||
|
||||
# target value
|
||||
discounted_q = rewards + gamma * next_q * mask
|
||||
|
||||
# Update critic
|
||||
loss_critic = torch.mean((current_q1 - discounted_q) ** 2) + torch.mean(
|
||||
(current_q2 - discounted_q) ** 2
|
||||
)
|
||||
|
||||
return loss_critic
|
||||
|
||||
def update_target_critic(self, tau):
|
||||
soft_update(self.target_q, self.critic_q, tau)
|
||||
@@ -0,0 +1,153 @@
|
||||
"""
|
||||
Reward-weighted regression (RWR) for diffusion policy.
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
import logging
|
||||
import einops
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
import torch.nn.functional as F
|
||||
|
||||
from model.diffusion.diffusion import DiffusionModel
|
||||
from model.diffusion.sampling import make_timesteps, extract
|
||||
|
||||
|
||||
class RWRDiffusion(DiffusionModel):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
use_ddim=False,
|
||||
# various clipping
|
||||
randn_clip_value=10,
|
||||
clamp_action=None,
|
||||
min_sampling_denoising_std=0.1,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(use_ddim=use_ddim, **kwargs)
|
||||
assert not self.use_ddim, "RWR does not support DDIM"
|
||||
|
||||
# For each denoising step, we clip sampled randn (from standard deviation) such that the sampled action is not too far away from mean
|
||||
self.randn_clip_value = randn_clip_value
|
||||
|
||||
# Action clamp range
|
||||
self.clamp_action = clamp_action
|
||||
|
||||
# Minimum std used in denoising process when sampling action - helps exploration
|
||||
self.min_sampling_denoising_std = min_sampling_denoising_std
|
||||
|
||||
# ---------- RL training ----------#
|
||||
|
||||
# override
|
||||
def p_losses(
|
||||
self,
|
||||
x_start,
|
||||
obs_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)
|
||||
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")
|
||||
loss *= rewards
|
||||
return loss.mean()
|
||||
|
||||
# ---------- Sampling ----------#
|
||||
|
||||
# override
|
||||
def p_mean_var(
|
||||
self,
|
||||
x,
|
||||
t,
|
||||
cond=None,
|
||||
):
|
||||
noise = self.network(x, t, cond=cond)
|
||||
|
||||
# Predict x_0
|
||||
if self.predict_epsilon:
|
||||
"""
|
||||
x₀ = √ 1\α̅ₜ xₜ - √ 1\α̅ₜ-1 ε
|
||||
"""
|
||||
x_recon = (
|
||||
extract(self.sqrt_recip_alphas_cumprod, t, x.shape) * x
|
||||
- extract(self.sqrt_recipm1_alphas_cumprod, t, x.shape) * noise
|
||||
)
|
||||
else: # directly predicting x₀
|
||||
x_recon = noise
|
||||
if self.denoised_clip_value is not None:
|
||||
x_recon.clamp_(-self.denoised_clip_value, self.denoised_clip_value)
|
||||
|
||||
# Get mu
|
||||
"""
|
||||
μₜ = β̃ₜ √ α̅ₜ₋₁/(1-α̅ₜ)x₀ + √ αₜ (1-α̅ₜ₋₁)/(1-α̅ₜ)xₜ
|
||||
"""
|
||||
mu = (
|
||||
extract(self.ddpm_mu_coef1, t, x.shape) * x_recon
|
||||
+ extract(self.ddpm_mu_coef2, t, x.shape) * x
|
||||
)
|
||||
logvar = extract(self.ddpm_logvar_clipped, t, x.shape)
|
||||
return mu, logvar
|
||||
|
||||
# override
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
cond,
|
||||
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]
|
||||
|
||||
# Loop
|
||||
x = torch.randn((B, self.horizon_steps, self.transition_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)
|
||||
mean, logvar = self.p_mean_var(
|
||||
x=x,
|
||||
t=t_b,
|
||||
cond=cond,
|
||||
)
|
||||
std = torch.exp(0.5 * logvar)
|
||||
|
||||
# Determine noise level
|
||||
if deterministic and t == 0:
|
||||
std = torch.zeros_like(std)
|
||||
elif deterministic: # For DDPM, sample with noise
|
||||
std = torch.clip(std, min=1e-3)
|
||||
else:
|
||||
std = torch.clip(std, min=self.min_sampling_denoising_std)
|
||||
noise = torch.randn_like(x).clamp_(
|
||||
-self.randn_clip_value, self.randn_clip_value
|
||||
)
|
||||
x = mean + std * noise
|
||||
|
||||
# clamp action at final step
|
||||
if self.clamp_action is not None and i == len(t_all) - 1:
|
||||
x = torch.clamp(x, -self.clamp_action, self.clamp_action)
|
||||
return x
|
||||
@@ -0,0 +1,465 @@
|
||||
"""
|
||||
Policy gradient with diffusion policy.
|
||||
|
||||
VPG: vanilla policy gradient
|
||||
|
||||
"""
|
||||
|
||||
import copy
|
||||
import torch
|
||||
import logging
|
||||
import einops
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
import torch.nn.functional as F
|
||||
|
||||
from model.diffusion.diffusion import DiffusionModel, Sample
|
||||
from model.diffusion.sampling import make_timesteps, extract
|
||||
from torch.distributions import Normal
|
||||
|
||||
|
||||
class VPGDiffusion(DiffusionModel):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
actor,
|
||||
critic,
|
||||
ft_denoising_steps,
|
||||
ft_denoising_steps_d=0,
|
||||
ft_denoising_steps_t=0,
|
||||
network_path=None,
|
||||
# various clipping
|
||||
randn_clip_value=10,
|
||||
clamp_action=False,
|
||||
min_sampling_denoising_std=0.1,
|
||||
min_logprob_denoising_std=0.1, # or the scheduler class
|
||||
eps_clip_value=None, # only used with DDIM
|
||||
# DDIM related
|
||||
eta=None,
|
||||
learn_eta=False,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(
|
||||
network=actor,
|
||||
network_path=network_path,
|
||||
**kwargs,
|
||||
)
|
||||
assert ft_denoising_steps <= self.denoising_steps
|
||||
assert ft_denoising_steps <= self.ddim_steps if self.use_ddim else True
|
||||
assert not (learn_eta and not self.use_ddim), "Cannot learn eta with DDPM."
|
||||
|
||||
# Number of denoising steps to use with fine-tuned model. Thus denoising_step - ft_denoising_steps is the number of denoising steps to use with original model.
|
||||
self.ft_denoising_steps = ft_denoising_steps
|
||||
self.ft_denoising_steps_d = ft_denoising_steps_d # annealing step size
|
||||
self.ft_denoising_steps_t = ft_denoising_steps_t # annealing interval
|
||||
self.ft_denoising_steps_cnt = 0
|
||||
|
||||
# Clip noise for numerical stability in policy gradient
|
||||
self.eps_clip_value = eps_clip_value
|
||||
|
||||
# For each denoising step, we clip sampled randn (from standard deviation) such that the sampled action is not too far away from mean
|
||||
self.randn_clip_value = randn_clip_value
|
||||
|
||||
# Whether to clamp sampled action between [-1, 1]
|
||||
self.clamp_action = clamp_action
|
||||
|
||||
# Minimum std used in denoising process when sampling action - helps exploration
|
||||
self.min_sampling_denoising_std = min_sampling_denoising_std
|
||||
|
||||
# Minimum std used in calculating denoising logprobs - for stability
|
||||
self.min_logprob_denoising_std = min_logprob_denoising_std
|
||||
|
||||
# Learnable eta
|
||||
self.learn_eta = learn_eta
|
||||
if eta is not None:
|
||||
self.eta = eta.to(self.device)
|
||||
if not learn_eta:
|
||||
for param in self.eta.parameters():
|
||||
param.requires_grad = False
|
||||
logging.info("Turned off gradients for eta")
|
||||
|
||||
# Re-name network to actor
|
||||
self.actor = self.network
|
||||
|
||||
# Make a copy of the original model
|
||||
self.actor_ft = copy.deepcopy(self.actor)
|
||||
logging.info("Cloned model for fine-tuning")
|
||||
|
||||
# Turn off gradients for original model
|
||||
for param in self.actor.parameters():
|
||||
param.requires_grad = False
|
||||
logging.info("Turned off gradients of the pretrained network")
|
||||
logging.info(
|
||||
f"Number of finetuned parameters: {sum(p.numel() for p in self.actor_ft.parameters() if p.requires_grad)}"
|
||||
)
|
||||
|
||||
# Value function
|
||||
self.critic = critic.to(self.device)
|
||||
if network_path is not None:
|
||||
checkpoint = torch.load(
|
||||
network_path, map_location=self.device, weights_only=True
|
||||
)
|
||||
if "ema" not in checkpoint: # load trained RL model
|
||||
self.load_state_dict(checkpoint["model"], strict=False)
|
||||
logging.info("Loaded critic from %s", network_path)
|
||||
|
||||
# ---------- 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
|
||||
if type(self.min_sampling_denoising_std) is not float:
|
||||
self.min_sampling_denoising_std.step()
|
||||
|
||||
# anneal denoising steps
|
||||
self.ft_denoising_steps_cnt += 1
|
||||
if (
|
||||
self.ft_denoising_steps_d > 0
|
||||
and self.ft_denoising_steps_t > 0
|
||||
and self.ft_denoising_steps_cnt % self.ft_denoising_steps_t == 0
|
||||
):
|
||||
self.ft_denoising_steps = max(
|
||||
0, self.ft_denoising_steps - self.ft_denoising_steps_d
|
||||
)
|
||||
|
||||
# update actor
|
||||
self.actor = self.actor_ft
|
||||
self.actor_ft = copy.deepcopy(self.actor)
|
||||
for param in self.actor.parameters():
|
||||
param.requires_grad = False
|
||||
logging.info(
|
||||
f"Finished annealing fine-tuning denoising steps to {self.ft_denoising_steps}"
|
||||
)
|
||||
|
||||
def get_min_sampling_denoising_std(self):
|
||||
if type(self.min_sampling_denoising_std) is float:
|
||||
return self.min_sampling_denoising_std
|
||||
else:
|
||||
return self.min_sampling_denoising_std()
|
||||
|
||||
# override
|
||||
def p_mean_var(
|
||||
self,
|
||||
x,
|
||||
t,
|
||||
cond=None,
|
||||
index=None,
|
||||
use_base_policy=False,
|
||||
deterministic=False,
|
||||
):
|
||||
noise = self.actor(x, t, cond=cond)
|
||||
if self.use_ddim:
|
||||
ft_indices = torch.where(
|
||||
index >= (self.ddim_steps - self.ft_denoising_steps)
|
||||
)[0]
|
||||
else:
|
||||
ft_indices = torch.where(t < self.ft_denoising_steps)[0]
|
||||
|
||||
# Use base policy to query expert model, e.g. for imitation loss
|
||||
actor = self.actor if use_base_policy else self.actor_ft
|
||||
|
||||
# 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)
|
||||
noise[ft_indices] = noise_ft
|
||||
|
||||
# Predict x_0
|
||||
if self.predict_epsilon:
|
||||
if self.use_ddim:
|
||||
"""
|
||||
x₀ = (xₜ - √ (1-αₜ) ε )/ √ αₜ
|
||||
"""
|
||||
alpha = extract(self.ddim_alphas, index, x.shape)
|
||||
alpha_prev = extract(self.ddim_alphas_prev, index, x.shape)
|
||||
sqrt_one_minus_alpha = extract(
|
||||
self.ddim_sqrt_one_minus_alphas, index, x.shape
|
||||
)
|
||||
x_recon = (x - sqrt_one_minus_alpha * noise) / (alpha**0.5)
|
||||
else:
|
||||
"""
|
||||
x₀ = √ 1\α̅ₜ xₜ - √ 1\α̅ₜ-1 ε
|
||||
"""
|
||||
x_recon = (
|
||||
extract(self.sqrt_recip_alphas_cumprod, t, x.shape) * x
|
||||
- extract(self.sqrt_recipm1_alphas_cumprod, t, x.shape) * noise
|
||||
)
|
||||
else: # directly predicting x₀
|
||||
x_recon = noise
|
||||
if self.denoised_clip_value is not None:
|
||||
x_recon.clamp_(-self.denoised_clip_value, self.denoised_clip_value)
|
||||
if self.use_ddim:
|
||||
# re-calculate noise based on clamped x_recon - default to false in HF, but let's use it here
|
||||
noise = (x - alpha ** (0.5) * x_recon) / sqrt_one_minus_alpha
|
||||
|
||||
# Clip epsilon for numerical stability in policy gradient - not sure if this is helpful yet, but the value can be huge sometimes. This has no effect if DDPM is used
|
||||
if self.use_ddim and self.eps_clip_value is not None:
|
||||
noise.clamp_(-self.eps_clip_value, self.eps_clip_value)
|
||||
|
||||
# Get mu
|
||||
if self.use_ddim:
|
||||
"""
|
||||
μ = √ αₜ₋₁ x₀ + √(1-αₜ₋₁ - σₜ²) ε
|
||||
"""
|
||||
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)
|
||||
sigma = (
|
||||
etas
|
||||
* ((1 - alpha_prev) / (1 - alpha) * (1 - alpha / alpha_prev)) ** 0.5
|
||||
).clamp_(min=1e-10)
|
||||
dir_xt_coef = (1.0 - alpha_prev - sigma**2).clamp_(min=0).sqrt()
|
||||
mu = (alpha_prev**0.5) * x_recon + dir_xt_coef * noise
|
||||
var = sigma**2
|
||||
logvar = torch.log(var)
|
||||
else:
|
||||
"""
|
||||
μₜ = β̃ₜ √ α̅ₜ₋₁/(1-α̅ₜ)x₀ + √ αₜ (1-α̅ₜ₋₁)/(1-α̅ₜ)xₜ
|
||||
"""
|
||||
mu = (
|
||||
extract(self.ddpm_mu_coef1, t, x.shape) * x_recon
|
||||
+ extract(self.ddpm_mu_coef2, t, x.shape) * x
|
||||
)
|
||||
logvar = extract(self.ddpm_logvar_clipped, t, x.shape)
|
||||
etas = torch.ones_like(mu).to(mu.device) # always one for DDPM
|
||||
return mu, logvar, etas
|
||||
|
||||
# override
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
cond,
|
||||
deterministic=False,
|
||||
return_chain=True,
|
||||
use_base_policy=False,
|
||||
):
|
||||
"""
|
||||
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
|
||||
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["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]
|
||||
|
||||
# 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)
|
||||
if self.use_ddim:
|
||||
t_all = self.ddim_t
|
||||
else:
|
||||
t_all = list(reversed(range(self.denoising_steps)))
|
||||
chain = [] if return_chain else None
|
||||
if not self.use_ddim and self.ft_denoising_steps == self.denoising_steps:
|
||||
chain.append(x)
|
||||
if self.use_ddim and self.ft_denoising_steps == self.ddim_steps:
|
||||
chain.append(x)
|
||||
for i, t in enumerate(t_all):
|
||||
t_b = make_timesteps(B, t, device)
|
||||
index_b = make_timesteps(B, i, device)
|
||||
mean, logvar, _ = self.p_mean_var(
|
||||
x=x,
|
||||
t=t_b,
|
||||
cond=cond,
|
||||
index=index_b,
|
||||
use_base_policy=use_base_policy,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
std = torch.exp(0.5 * logvar)
|
||||
|
||||
# Determine noise level
|
||||
if self.use_ddim:
|
||||
if deterministic:
|
||||
std = torch.zeros_like(std)
|
||||
else:
|
||||
std = torch.clip(std, min=min_sampling_denoising_std)
|
||||
else:
|
||||
if deterministic and t == 0:
|
||||
std = torch.zeros_like(std)
|
||||
elif deterministic: # still keep the original noise
|
||||
std = torch.clip(std, min=1e-3)
|
||||
else: # use higher minimum noise
|
||||
std = torch.clip(std, min=min_sampling_denoising_std)
|
||||
noise = torch.randn_like(x).clamp_(
|
||||
-self.randn_clip_value, self.randn_clip_value
|
||||
)
|
||||
x = mean + std * noise
|
||||
|
||||
# clamp action at final step
|
||||
if self.clamp_action and i == len(t_all) - 1:
|
||||
x = torch.clamp(x, -1, 1)
|
||||
|
||||
if return_chain:
|
||||
if not self.use_ddim and t <= self.ft_denoising_steps:
|
||||
chain.append(x)
|
||||
elif self.use_ddim and i >= (
|
||||
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)
|
||||
|
||||
# ---------- RL training ----------#
|
||||
|
||||
def get_logprobs(
|
||||
self,
|
||||
obs,
|
||||
chains,
|
||||
get_ent: bool = False,
|
||||
use_base_policy: bool = False,
|
||||
):
|
||||
"""
|
||||
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
|
||||
|
||||
Returns:
|
||||
logprobs: (B x num_denoising_steps, horizon_step, action_dim)
|
||||
entropy (if get_ent=True): (B x num_denoising_steps, horizon_step)
|
||||
"""
|
||||
# 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 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.repeat(chains.shape[0], 1).flatten()
|
||||
if self.use_ddim:
|
||||
indices_single = torch.arange(
|
||||
start=self.ddim_steps - self.ft_denoising_steps,
|
||||
end=self.ddim_steps,
|
||||
device=self.device,
|
||||
) # only used for DDIM
|
||||
indices = indices_single.repeat(chains.shape[0])
|
||||
else:
|
||||
indices = None
|
||||
|
||||
# Split chains
|
||||
chains_prev = chains[:, :-1]
|
||||
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)
|
||||
|
||||
# Forward pass with previous chains
|
||||
next_mean, logvar, eta = self.p_mean_var(
|
||||
chains_prev,
|
||||
t_all,
|
||||
cond=cond,
|
||||
index=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, obs, 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)
|
||||
"""
|
||||
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()
|
||||
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)
|
||||
|
||||
# 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 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
|
||||
|
||||
# Sum/avg over horizon steps
|
||||
logprobs = logprobs.mean(-1) # -> (n_steps x n_envs)
|
||||
|
||||
# Get REINFORCE loss
|
||||
loss_actor = torch.mean(-logprobs * advantage)
|
||||
|
||||
# Train critic to predict state value
|
||||
pred = self.critic(obs).squeeze()
|
||||
loss_critic = F.mse_loss(pred, reward)
|
||||
return loss_actor, loss_critic, eta
|
||||
@@ -0,0 +1,164 @@
|
||||
"""
|
||||
Eta in DDIM.
|
||||
|
||||
Can be learned but always fixed to 1 during training and 0 during eval right now.
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
from model.common.mlp import MLP
|
||||
|
||||
|
||||
class EtaFixed(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_eta=0.5,
|
||||
min_eta=0.1,
|
||||
max_eta=1.0,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.eta_logit = torch.nn.Parameter(torch.ones(1))
|
||||
self.min = min_eta
|
||||
self.max = max_eta
|
||||
|
||||
# initialize such that eta = base_eta
|
||||
self.eta_logit.data = torch.atanh(
|
||||
torch.tensor([2 * (base_eta - min_eta) / (max_eta - min_eta) - 1])
|
||||
)
|
||||
|
||||
def __call__(self, x):
|
||||
"""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
|
||||
eta_normalized = torch.tanh(self.eta_logit)
|
||||
|
||||
# map to min and max from [-1, 1]
|
||||
eta = 0.5 * (eta_normalized + 1) * (self.max - self.min) + self.min
|
||||
return torch.full((B, 1), eta.item()).to(device)
|
||||
|
||||
|
||||
class EtaAction(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
action_dim,
|
||||
base_eta=0.5,
|
||||
min_eta=0.1,
|
||||
max_eta=1.0,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
# initialize such that eta = base_eta
|
||||
self.eta_logit = torch.nn.Parameter(
|
||||
torch.ones(action_dim)
|
||||
* torch.atanh(
|
||||
torch.tensor([2 * (base_eta - min_eta) / (max_eta - min_eta) - 1])
|
||||
)
|
||||
)
|
||||
self.min = min_eta
|
||||
self.max = max_eta
|
||||
|
||||
def __call__(self, x):
|
||||
"""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
|
||||
eta_normalized = torch.tanh(self.eta_logit)
|
||||
|
||||
# map to min and max from [-1, 1]
|
||||
eta = 0.5 * (eta_normalized + 1) * (self.max - self.min) + self.min
|
||||
return eta.repeat(B, 1).to(device)
|
||||
|
||||
|
||||
class EtaState(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_dim,
|
||||
mlp_dims,
|
||||
activation_type="ReLU",
|
||||
out_activation_type="Identity",
|
||||
base_eta=0.5,
|
||||
min_eta=0.1,
|
||||
max_eta=1.0,
|
||||
gain=1e-2,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.base = base_eta
|
||||
self.min_res = min_eta - base_eta
|
||||
self.max_res = max_eta - base_eta
|
||||
self.mlp_res = MLP(
|
||||
[input_dim] + mlp_dims + [1],
|
||||
activation_type=activation_type,
|
||||
out_activation_type=out_activation_type,
|
||||
)
|
||||
|
||||
# initialize such that mlp(x) = 0
|
||||
for m in self.mlp_res.modules():
|
||||
if isinstance(m, torch.nn.Linear):
|
||||
torch.nn.init.xavier_normal_(m.weight, gain=gain)
|
||||
m.bias.data.fill_(0)
|
||||
|
||||
def __call__(self, x):
|
||||
if isinstance(x, dict):
|
||||
raise NotImplementedError(
|
||||
"State-based eta not implemented for image-based training!"
|
||||
)
|
||||
x = x.view(x.size(0), -1)
|
||||
eta_res = self.mlp_res(x)
|
||||
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)
|
||||
|
||||
|
||||
class EtaStateAction(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_dim,
|
||||
mlp_dims,
|
||||
action_dim,
|
||||
activation_type="ReLU",
|
||||
out_activation_type="Identity",
|
||||
base_eta=1,
|
||||
min_eta=1e-3,
|
||||
max_eta=2,
|
||||
gain=1e-2,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.base = base_eta
|
||||
self.min_res = min_eta - base_eta
|
||||
self.max_res = max_eta - base_eta
|
||||
self.mlp_res = MLP(
|
||||
[input_dim] + mlp_dims + [action_dim],
|
||||
activation_type=activation_type,
|
||||
out_activation_type=out_activation_type,
|
||||
)
|
||||
|
||||
# initialize such that mlp(x) = 0
|
||||
for m in self.mlp_res.modules():
|
||||
if isinstance(m, torch.nn.Linear):
|
||||
torch.nn.init.xavier_normal_(m.weight, gain=gain)
|
||||
m.bias.data.fill_(0)
|
||||
|
||||
def __call__(self, x):
|
||||
if isinstance(x, dict):
|
||||
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)
|
||||
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)
|
||||
@@ -0,0 +1,188 @@
|
||||
"""
|
||||
Solving probabilistic ODE for exact likelihood, from https://github.com/yang-song/score_sde_pytorch
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
from torchdiffeq import odeint
|
||||
|
||||
# adjoint can reduce memory, but not faster
|
||||
# from torchdiffeq import odeint_adjoint as odeint
|
||||
from model.diffusion.sde_lib import get_score_fn
|
||||
|
||||
|
||||
def get_likelihood_fn(
|
||||
sde,
|
||||
hutchinson_type="Rademacher",
|
||||
rtol=1e-5,
|
||||
atol=1e-5,
|
||||
method="RK45",
|
||||
steps=10, # should not matter, only t_eval
|
||||
step_size=1e-3,
|
||||
eps=1e-5,
|
||||
continuous=False,
|
||||
probability_flow=False,
|
||||
predict_epsilon=False,
|
||||
num_epsilon=1,
|
||||
):
|
||||
"""Create a function to compute the unbiased log-likelihood estimate of a given data point.
|
||||
|
||||
Args:
|
||||
sde: A `sde_lib.SDE` object that represents the forward SDE.
|
||||
inverse_scaler: The inverse data normalizer.
|
||||
hutchinson_type: "Rademacher" or "Gaussian". The type of noise for Hutchinson-Skilling trace estimator.
|
||||
rtol: A `float` number. The relative tolerance level of the black-box ODE solver.
|
||||
atol: A `float` number. The absolute tolerance level of the black-box ODE solver.
|
||||
method: A `str`. The algorithm for the black-box ODE solver.
|
||||
See documentation for `scipy.integrate.solve_ivp`.
|
||||
eps: A `float` number. The probability flow ODE is integrated to `eps` for numerical stability.
|
||||
|
||||
Returns:
|
||||
A function that a batch of data points and returns the log-likelihoods in bits/dim,
|
||||
the latent code, and the number of function evaluations cost by computation.
|
||||
"""
|
||||
|
||||
def drift_fn(
|
||||
model,
|
||||
x,
|
||||
t,
|
||||
**kwargs,
|
||||
):
|
||||
"""The drift function of the reverse-time SDE."""
|
||||
score_fn = get_score_fn(
|
||||
sde,
|
||||
model,
|
||||
continuous=continuous,
|
||||
predict_epsilon=predict_epsilon,
|
||||
)
|
||||
# Probability flow ODE is a special case of Reverse SDE
|
||||
rsde = sde.reverse(score_fn, probability_flow=probability_flow)
|
||||
sde_out = rsde.sde(x, t, **kwargs)[0]
|
||||
return sde_out
|
||||
|
||||
def div_fn(
|
||||
model,
|
||||
x,
|
||||
t,
|
||||
noise,
|
||||
create_graph=False,
|
||||
**kwargs,
|
||||
):
|
||||
with torch.enable_grad():
|
||||
x.requires_grad_(True)
|
||||
fn_eps = torch.sum(drift_fn(model, x, t, **kwargs) * noise)
|
||||
grad_fn_eps = torch.autograd.grad(
|
||||
fn_eps,
|
||||
x,
|
||||
create_graph=create_graph,
|
||||
)[0]
|
||||
if not create_graph:
|
||||
x.requires_grad_(False)
|
||||
return torch.sum(grad_fn_eps * noise, dim=(1, 2))
|
||||
|
||||
def likelihood_fn(
|
||||
model,
|
||||
model_ft,
|
||||
data,
|
||||
denoising_steps,
|
||||
ft_denoising_steps,
|
||||
cond,
|
||||
**kwargs,
|
||||
):
|
||||
"""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
|
||||
|
||||
Returns:
|
||||
logprob: B
|
||||
"""
|
||||
shape = data.shape
|
||||
B, H, A = shape
|
||||
device = data.device
|
||||
|
||||
# sample epsilon
|
||||
if hutchinson_type == "Gaussian":
|
||||
epsilon = torch.randn(size=(B * num_epsilon, H, A), device=device)
|
||||
elif hutchinson_type == "Rademacher":
|
||||
epsilon = (
|
||||
torch.randint(
|
||||
low=0, high=2, size=(B * num_epsilon, H, A), device=device
|
||||
).float()
|
||||
* 2
|
||||
- 1.0
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"Hutchinson type {hutchinson_type} unknown.")
|
||||
|
||||
# repeat for expectation
|
||||
cond_eps = cond.repeat_interleave(num_epsilon, dim=0)
|
||||
|
||||
def ode_func(t, x):
|
||||
x = x[:, :-1]
|
||||
vec_t = torch.full(
|
||||
(x.shape[0],),
|
||||
torch.round(t * denoising_steps),
|
||||
device=x.device,
|
||||
dtype=int,
|
||||
)
|
||||
if torch.round(t * denoising_steps) <= ft_denoising_steps:
|
||||
model_fn = model_ft
|
||||
else:
|
||||
model_fn = model
|
||||
x = x.view(shape) # B x horizon x transition_dim
|
||||
drift = drift_fn(
|
||||
model_fn,
|
||||
x,
|
||||
vec_t,
|
||||
cond=cond,
|
||||
**kwargs,
|
||||
).reshape(B, -1)
|
||||
|
||||
# repeat for expectation
|
||||
x = x.repeat_interleave(num_epsilon, dim=0)
|
||||
vec_t = vec_t.repeat_interleave(num_epsilon)
|
||||
|
||||
logp_grad = div_fn(
|
||||
model,
|
||||
x,
|
||||
vec_t,
|
||||
epsilon,
|
||||
create_graph=True,
|
||||
cond=cond_eps,
|
||||
**kwargs,
|
||||
)[:, None].reshape(B, num_epsilon, -1)
|
||||
logp_grad = logp_grad.mean(dim=1) # expectation over epsilon
|
||||
return torch.cat(
|
||||
[drift, logp_grad], dim=-1
|
||||
) # Concatenate along the feature dimension
|
||||
|
||||
# flatten data
|
||||
data = data.view(shape[0], -1)
|
||||
init = torch.hstack(
|
||||
(data, torch.zeros((shape[0], 1)).to(data.dtype).to(device))
|
||||
)
|
||||
t_eval = torch.linspace(eps, sde.T, steps=steps).to(device) # eval points
|
||||
solution = odeint(
|
||||
ode_func,
|
||||
init,
|
||||
t_eval,
|
||||
method=method,
|
||||
rtol=rtol,
|
||||
atol=atol,
|
||||
options={"step_size": step_size},
|
||||
# args=(model, epsilon),
|
||||
) # steps x batch x 3
|
||||
zp = solution[-1] # batch x 3
|
||||
z = zp[:, :-1].view(shape)
|
||||
delta_logp = zp[:, -1]
|
||||
prior_logp = sde.prior_logp(z)
|
||||
N = torch.prod(torch.tensor(shape[1:]))
|
||||
# print("prior:", prior_logp / (np.log(2) * N))
|
||||
# print("delta:", delta_logp / (np.log(2) * N))
|
||||
logprob = (prior_logp + delta_logp) / (np.log(2) * N)
|
||||
return logprob
|
||||
|
||||
return likelihood_fn
|
||||
@@ -0,0 +1,233 @@
|
||||
"""
|
||||
MLP models for diffusion policies.
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import logging
|
||||
import einops
|
||||
from copy import deepcopy
|
||||
|
||||
from model.common.mlp import MLP, ResidualMLP
|
||||
from model.diffusion.modules import SinusoidalPosEmb
|
||||
from model.common.modules import SpatialEmb, RandomShiftsAug
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class VisionDiffusionMLP(nn.Module):
|
||||
"""With ViT backbone"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
backbone,
|
||||
transition_dim,
|
||||
horizon_steps,
|
||||
cond_dim,
|
||||
time_dim=16,
|
||||
mlp_dims=[256, 256],
|
||||
activation_type="Mish",
|
||||
out_activation_type="Identity",
|
||||
use_layernorm=False,
|
||||
residual_style=False,
|
||||
spatial_emb=0,
|
||||
visual_feature_dim=128,
|
||||
repr_dim=96 * 96,
|
||||
patch_repr_dim=128,
|
||||
dropout=0,
|
||||
num_img=1,
|
||||
augment=False,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# vision
|
||||
self.backbone = backbone
|
||||
if augment:
|
||||
self.aug = RandomShiftsAug(pad=4)
|
||||
self.augment = augment
|
||||
if spatial_emb > 0:
|
||||
assert spatial_emb > 1, "this is the dimension"
|
||||
if num_img > 1:
|
||||
self.compress1 = SpatialEmb(
|
||||
num_patch=121, # TODO: repr_dim // patch_repr_dim,
|
||||
patch_dim=patch_repr_dim,
|
||||
prop_dim=cond_dim,
|
||||
proj_dim=spatial_emb,
|
||||
dropout=dropout,
|
||||
)
|
||||
self.compress2 = deepcopy(self.compress1)
|
||||
else: # TODO: clean up
|
||||
self.compress = SpatialEmb(
|
||||
num_patch=121,
|
||||
patch_dim=patch_repr_dim,
|
||||
prop_dim=cond_dim,
|
||||
proj_dim=spatial_emb,
|
||||
dropout=dropout,
|
||||
)
|
||||
visual_feature_dim = spatial_emb * num_img
|
||||
else:
|
||||
self.compress = nn.Sequential(
|
||||
nn.Linear(repr_dim, visual_feature_dim),
|
||||
nn.LayerNorm(visual_feature_dim),
|
||||
nn.Dropout(dropout),
|
||||
nn.ReLU(),
|
||||
)
|
||||
|
||||
# diffusion
|
||||
input_dim = (
|
||||
time_dim + transition_dim * horizon_steps + visual_feature_dim + cond_dim
|
||||
)
|
||||
output_dim = transition_dim * horizon_steps
|
||||
self.time_embedding = nn.Sequential(
|
||||
SinusoidalPosEmb(time_dim),
|
||||
nn.Linear(time_dim, time_dim * 2),
|
||||
nn.Mish(),
|
||||
nn.Linear(time_dim * 2, time_dim),
|
||||
)
|
||||
if residual_style:
|
||||
model = ResidualMLP
|
||||
else:
|
||||
model = MLP
|
||||
self.mlp_mean = model(
|
||||
[input_dim] + mlp_dims + [output_dim],
|
||||
activation_type=activation_type,
|
||||
out_activation_type=out_activation_type,
|
||||
use_layernorm=use_layernorm,
|
||||
)
|
||||
self.time_dim = time_dim
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
time,
|
||||
cond=None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
x: (B,T,obs_dim)
|
||||
time: (B,) or int, diffusion step
|
||||
cond: dict (B,cond_step,cond_dim)
|
||||
output: (B,T,input_dim)
|
||||
"""
|
||||
# flatten T and input_dim
|
||||
B, T, input_dim = x.shape
|
||||
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")
|
||||
else:
|
||||
rgb = cond["rgb"]
|
||||
if cond["state"].ndim == 3:
|
||||
state = einops.rearrange(cond["state"], "b d c -> (b d) c")
|
||||
else:
|
||||
state = cond["state"]
|
||||
|
||||
# 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.augment:
|
||||
rgb1 = self.aug(rgb1)
|
||||
rgb2 = self.aug(rgb2)
|
||||
feat1 = self.backbone(rgb1)
|
||||
feat2 = self.backbone(rgb2)
|
||||
feat1 = self.compress1.forward(feat1, state)
|
||||
feat2 = self.compress2.forward(feat2, state)
|
||||
feat = torch.cat([feat1, feat2], dim=-1)
|
||||
else: # single image
|
||||
if self.augment:
|
||||
rgb = self.aug(rgb) # uint8 -> float32
|
||||
feat = self.backbone(rgb)
|
||||
|
||||
# compress
|
||||
if isinstance(self.compress, SpatialEmb):
|
||||
feat = self.compress.forward(feat, state)
|
||||
else:
|
||||
feat = feat.flatten(1, -1)
|
||||
feat = self.compress(feat)
|
||||
cond_encoded = torch.cat([feat, state], dim=-1)
|
||||
|
||||
# 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_encoded], dim=-1)
|
||||
|
||||
# mlp
|
||||
out = self.mlp_mean(x)
|
||||
return out.view(B, T, input_dim)
|
||||
|
||||
|
||||
class DiffusionMLP(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transition_dim,
|
||||
horizon_steps,
|
||||
cond_dim,
|
||||
time_dim=16,
|
||||
mlp_dims=[256, 256],
|
||||
cond_mlp_dims=None,
|
||||
activation_type="Mish",
|
||||
out_activation_type="Identity",
|
||||
use_layernorm=False,
|
||||
residual_style=False,
|
||||
):
|
||||
super().__init__()
|
||||
output_dim = transition_dim * horizon_steps
|
||||
self.time_embedding = nn.Sequential(
|
||||
SinusoidalPosEmb(time_dim),
|
||||
nn.Linear(time_dim, time_dim * 2),
|
||||
nn.Mish(),
|
||||
nn.Linear(time_dim * 2, time_dim),
|
||||
)
|
||||
if residual_style:
|
||||
model = ResidualMLP
|
||||
else:
|
||||
model = MLP
|
||||
if cond_mlp_dims is not None:
|
||||
self.cond_mlp = MLP(
|
||||
[cond_dim] + cond_mlp_dims,
|
||||
activation_type=activation_type,
|
||||
out_activation_type="Identity",
|
||||
)
|
||||
input_dim = time_dim + transition_dim * horizon_steps + cond_mlp_dims[-1]
|
||||
else:
|
||||
input_dim = time_dim + transition_dim * horizon_steps + cond_dim
|
||||
self.mlp_mean = model(
|
||||
[input_dim] + mlp_dims + [output_dim],
|
||||
activation_type=activation_type,
|
||||
out_activation_type=out_activation_type,
|
||||
use_layernorm=use_layernorm,
|
||||
)
|
||||
self.time_dim = time_dim
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
time,
|
||||
cond=None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
x: (B,T,obs_dim)
|
||||
time: (B,) or int, diffusion step
|
||||
cond: (B,cond_step,cond_dim)
|
||||
output: (B,T,input_dim)
|
||||
"""
|
||||
# flatten T and input_dim
|
||||
B, T, input_dim = x.shape
|
||||
x = x.view(B, -1)
|
||||
cond = cond.view(B, -1) if cond is not None else None
|
||||
if hasattr(self, "cond_mlp"):
|
||||
cond = self.cond_mlp(cond)
|
||||
|
||||
# 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)
|
||||
|
||||
# mlp
|
||||
out = self.mlp_mean(x)
|
||||
return out.view(B, T, input_dim)
|
||||
@@ -0,0 +1,95 @@
|
||||
"""
|
||||
From Diffuser https://github.com/jannerm/diffuser
|
||||
|
||||
For MLP and UNet diffusion models.
|
||||
|
||||
"""
|
||||
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops.layers.torch import Rearrange
|
||||
|
||||
|
||||
class SinusoidalPosEmb(nn.Module):
|
||||
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
|
||||
def forward(self, x):
|
||||
device = x.device
|
||||
half_dim = self.dim // 2
|
||||
emb = math.log(10000) / (half_dim - 1)
|
||||
emb = torch.exp(torch.arange(half_dim, device=device) * -emb)
|
||||
emb = x[:, None] * emb[None, :]
|
||||
emb = torch.cat((emb.sin(), emb.cos()), dim=-1)
|
||||
return emb
|
||||
|
||||
|
||||
class Downsample1d(nn.Module):
|
||||
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
self.conv = nn.Conv1d(dim, dim, 3, 2, 1)
|
||||
|
||||
def forward(self, x):
|
||||
return self.conv(x)
|
||||
|
||||
|
||||
class Upsample1d(nn.Module):
|
||||
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
self.conv = nn.ConvTranspose1d(dim, dim, 4, 2, 1)
|
||||
|
||||
def forward(self, x):
|
||||
return self.conv(x)
|
||||
|
||||
|
||||
class Conv1dBlock(nn.Module):
|
||||
"""
|
||||
Conv1d --> GroupNorm --> Mish
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
inp_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
n_groups=None,
|
||||
activation_type="Mish",
|
||||
eps=1e-5,
|
||||
):
|
||||
super().__init__()
|
||||
if activation_type == "Mish":
|
||||
act = nn.Mish()
|
||||
elif activation_type == "ReLU":
|
||||
act = nn.ReLU()
|
||||
else:
|
||||
raise "Unknown activation type for Conv1dBlock"
|
||||
|
||||
self.block = nn.Sequential(
|
||||
nn.Conv1d(
|
||||
inp_channels, out_channels, kernel_size, padding=kernel_size // 2
|
||||
),
|
||||
(
|
||||
Rearrange("batch channels horizon -> batch channels 1 horizon")
|
||||
if n_groups is not None
|
||||
else nn.Identity()
|
||||
),
|
||||
(
|
||||
nn.GroupNorm(n_groups, out_channels, eps=eps)
|
||||
if n_groups is not None
|
||||
else nn.Identity()
|
||||
),
|
||||
(
|
||||
Rearrange("batch channels 1 horizon -> batch channels horizon")
|
||||
if n_groups is not None
|
||||
else nn.Identity()
|
||||
),
|
||||
act,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
return self.block(x)
|
||||
@@ -0,0 +1,37 @@
|
||||
"""
|
||||
From Diffuser https://github.com/jannerm/diffuser
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
|
||||
def cosine_beta_schedule(timesteps, s=0.008, dtype=torch.float32):
|
||||
"""
|
||||
cosine schedule as proposed in https://openreview.net/forum?id=-NEXDKk8gZ
|
||||
"""
|
||||
steps = timesteps + 1
|
||||
x = np.linspace(0, steps, steps)
|
||||
alphas_cumprod = np.cos(((x / steps) + s) / (1 + s) * np.pi * 0.5) ** 2
|
||||
alphas_cumprod = alphas_cumprod / alphas_cumprod[0]
|
||||
betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])
|
||||
betas_clipped = np.clip(betas, a_min=0, a_max=0.999)
|
||||
return torch.tensor(betas_clipped, dtype=dtype)
|
||||
|
||||
|
||||
def extract(a, t, x_shape):
|
||||
b, *_ = t.shape
|
||||
out = a.gather(-1, t)
|
||||
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
|
||||
@@ -0,0 +1,213 @@
|
||||
"""
|
||||
Abstract SDE classes, Reverse SDE, and VE/VP SDEs.
|
||||
|
||||
From https://github.com/yang-song/score_sde_pytorch
|
||||
|
||||
"""
|
||||
|
||||
import abc
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
|
||||
def get_score_fn(
|
||||
sde,
|
||||
model,
|
||||
continuous=False,
|
||||
predict_epsilon=False,
|
||||
):
|
||||
"""Wraps `score_fn` so that the model output corresponds to a real time-dependent score function.
|
||||
|
||||
Args:
|
||||
sde: An `sde_lib.SDE` object that represents the forward SDE.
|
||||
model: A score model.
|
||||
continuous: If `True`, the score-based model is expected to directly take continuous time steps.
|
||||
|
||||
Returns:
|
||||
A score function.
|
||||
"""
|
||||
|
||||
def score_fn(x, t, **kwargs):
|
||||
"""
|
||||
Use [:, None, None] to add two dimensions (horizon and transition)
|
||||
"""
|
||||
score = model(x, t, **kwargs)
|
||||
|
||||
if not predict_epsilon: # get epsilon first from predicted mu
|
||||
score = (
|
||||
-(x - score * sde.sqrt_alphas[t.long()][:, None, None])
|
||||
/ sde.discrete_betas[t.long()][:, None, None]
|
||||
)
|
||||
else:
|
||||
std = sde.sqrt_1m_alpha_bar[t.long()]
|
||||
score = -score / std[:, None, None]
|
||||
return score
|
||||
|
||||
return score_fn
|
||||
|
||||
|
||||
class SDE(abc.ABC):
|
||||
"""SDE abstract class. Functions are designed for a mini-batch of inputs."""
|
||||
|
||||
def __init__(self, N):
|
||||
"""Construct an SDE.
|
||||
|
||||
Args:
|
||||
N: number of discretization time steps.
|
||||
"""
|
||||
super().__init__()
|
||||
self.N = N
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def T(self):
|
||||
"""End time of the SDE."""
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
def sde(self, x, t):
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
def marginal_prob(self, x, t):
|
||||
"""Parameters to determine the marginal distribution of the SDE, $p_t(x)$."""
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
def prior_sampling(self, shape):
|
||||
"""Generate one sample from the prior distribution, $p_T(x)$."""
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
def prior_logp(self, z):
|
||||
"""Compute log-density of the prior distribution.
|
||||
|
||||
Useful for computing the log-likelihood via probability flow ODE.
|
||||
|
||||
Args:
|
||||
z: latent code
|
||||
Returns:
|
||||
log probability density
|
||||
"""
|
||||
pass
|
||||
|
||||
def discretize(self, x, t):
|
||||
"""Discretize the SDE in the form: x_{i+1} = x_i + f_i(x_i) + G_i z_i.
|
||||
|
||||
Useful for reverse diffusion sampling and probabiliy flow sampling.
|
||||
Defaults to Euler-Maruyama discretization.
|
||||
|
||||
Args:
|
||||
x: a torch tensor
|
||||
t: a torch float representing the time step (from 0 to `self.T`)
|
||||
|
||||
Returns:
|
||||
f, G
|
||||
"""
|
||||
dt = 1 / self.N
|
||||
drift, diffusion = self.sde(x, t)
|
||||
f = drift * dt
|
||||
G = diffusion * torch.sqrt(torch.tensor(dt, device=t.device))
|
||||
return f, G
|
||||
|
||||
def reverse(self, score_fn, probability_flow=False):
|
||||
"""Create the reverse-time SDE/ODE.
|
||||
|
||||
Args:
|
||||
score_fn: A time-dependent score-based model that takes x and t and returns the score.
|
||||
probability_flow: If `True`, create the reverse-time ODE used for probability flow sampling.
|
||||
"""
|
||||
N = self.N
|
||||
T = self.T
|
||||
sde_fn = self.sde
|
||||
discretize_fn = self.discretize
|
||||
|
||||
# Build the class for reverse-time SDE.
|
||||
class RSDE(self.__class__):
|
||||
def __init__(self):
|
||||
self.N = N
|
||||
self.probability_flow = probability_flow
|
||||
|
||||
@property
|
||||
def T(self):
|
||||
return T
|
||||
|
||||
def sde(self, x, t, **kwargs):
|
||||
"""Create the drift and diffusion functions for the reverse SDE/ODE."""
|
||||
drift, diffusion = sde_fn(x, t)
|
||||
score = score_fn(x, t, **kwargs)
|
||||
drift = drift - diffusion[:, None, None] ** 2 * score * (
|
||||
0.5 if self.probability_flow else 1.0
|
||||
)
|
||||
# Set the diffusion function to zero for ODEs.
|
||||
diffusion = 0.0 if self.probability_flow else diffusion
|
||||
return drift, diffusion
|
||||
|
||||
def discretize(self, x, t):
|
||||
"""Create discretized iteration rules for the reverse diffusion sampler."""
|
||||
f, G = discretize_fn(x, t)
|
||||
rev_f = f - G[:, None] ** 2 * score_fn(x, t) * (
|
||||
0.5 if self.probability_flow else 1.0
|
||||
)
|
||||
rev_G = torch.zeros_like(G) if self.probability_flow else G
|
||||
return rev_f, rev_G
|
||||
|
||||
return RSDE()
|
||||
|
||||
|
||||
class VPSDE(SDE):
|
||||
def __init__(self, N=1000):
|
||||
"""Construct a Variance Preserving SDE.
|
||||
|
||||
Args:
|
||||
beta_min: value of beta(0)
|
||||
beta_max: value of beta(1)
|
||||
N: number of discretization steps
|
||||
"""
|
||||
super().__init__(N)
|
||||
|
||||
def set_betas(self, betas, min_beta=0.01):
|
||||
self.discrete_betas = betas.clamp(min=min_beta) # cosine schedule from our DDPM
|
||||
self.alphas = 1.0 - self.discrete_betas
|
||||
self.sqrt_alphas = torch.sqrt(self.alphas)
|
||||
self.alphas_bar = torch.cumprod(self.alphas, axis=0)
|
||||
self.sqrt_1m_alpha_bar = torch.sqrt(1 - self.alphas_bar)
|
||||
|
||||
@property
|
||||
def T(self):
|
||||
return 1
|
||||
|
||||
def sde(self, x, t):
|
||||
# dx = - 1/2 beta(t) x dt + sqrt(beta(t)) dW
|
||||
beta_t = self.discrete_betas[t]
|
||||
drift = -0.5 * beta_t[:, None, None] * x
|
||||
diffusion = torch.sqrt(beta_t)
|
||||
return drift, diffusion
|
||||
|
||||
def marginal_prob(self, x, t):
|
||||
raise NotImplementedError
|
||||
# log_mean_coeff = (
|
||||
# -0.25 * t**2 * (self.beta_1 - self.beta_0) - 0.5 * t * self.beta_0
|
||||
# )
|
||||
# mean = torch.exp(log_mean_coeff[:, None, None]) * x
|
||||
# std = torch.sqrt(1.0 - torch.exp(2.0 * log_mean_coeff))
|
||||
# return mean, std
|
||||
|
||||
def prior_sampling(self, shape):
|
||||
return torch.randn(*shape)
|
||||
|
||||
def prior_logp(self, z):
|
||||
shape = z.shape
|
||||
N = np.prod(shape[1:])
|
||||
logps = -N / 2.0 * np.log(2 * np.pi) - torch.sum(z**2, dim=(1, 2)) / 2.0
|
||||
return logps
|
||||
|
||||
def discretize(self, x, t):
|
||||
"""DDPM discretization."""
|
||||
timestep = (t * (self.N - 1) / self.T).long()
|
||||
beta = self.discrete_betas.to(x.device)[timestep]
|
||||
alpha = self.alphas.to(x.device)[timestep]
|
||||
sqrt_beta = torch.sqrt(beta)
|
||||
f = torch.sqrt(alpha)[:, None, None] * x - x
|
||||
G = sqrt_beta
|
||||
return f, G
|
||||
@@ -0,0 +1,318 @@
|
||||
"""
|
||||
UNet implementation. Minorly modified from Diffusion Policy: https://github.com/columbia-ai-robotics/diffusion_policy/blob/main/diffusion_policy/model/diffusion/conv1d_components.py
|
||||
|
||||
Set `smaller_encoder` to False for using larger observation encoder in ResidualBlock1D
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import einops
|
||||
from einops.layers.torch import Rearrange
|
||||
import logging
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
from model.diffusion.modules import (
|
||||
SinusoidalPosEmb,
|
||||
Downsample1d,
|
||||
Upsample1d,
|
||||
Conv1dBlock,
|
||||
)
|
||||
from model.common.mlp import ResidualMLP
|
||||
|
||||
|
||||
class ResidualBlock1D(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
cond_dim,
|
||||
kernel_size=5,
|
||||
n_groups=None,
|
||||
cond_predict_scale=False,
|
||||
larger_encoder=False,
|
||||
activation_type="Mish",
|
||||
groupnorm_eps=1e-5,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
Conv1dBlock(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
n_groups=n_groups,
|
||||
activation_type=activation_type,
|
||||
eps=groupnorm_eps,
|
||||
),
|
||||
Conv1dBlock(
|
||||
out_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
n_groups=n_groups,
|
||||
activation_type=activation_type,
|
||||
eps=groupnorm_eps,
|
||||
),
|
||||
]
|
||||
)
|
||||
if activation_type == "Mish":
|
||||
act = nn.Mish()
|
||||
elif activation_type == "ReLU":
|
||||
act = nn.ReLU()
|
||||
else:
|
||||
raise "Unknown activation type for ConditionalResidualBlock1D"
|
||||
|
||||
# FiLM modulation https://arxiv.org/abs/1709.07871
|
||||
# predicts per-channel scale and bias
|
||||
cond_channels = out_channels
|
||||
if cond_predict_scale:
|
||||
cond_channels = out_channels * 2
|
||||
self.cond_predict_scale = cond_predict_scale
|
||||
self.out_channels = out_channels
|
||||
if larger_encoder:
|
||||
self.cond_encoder = nn.Sequential(
|
||||
nn.Linear(cond_dim, cond_channels),
|
||||
act,
|
||||
nn.Linear(cond_channels, cond_channels),
|
||||
act,
|
||||
nn.Linear(cond_channels, cond_channels),
|
||||
Rearrange("batch t -> batch t 1"),
|
||||
)
|
||||
else:
|
||||
self.cond_encoder = nn.Sequential(
|
||||
act,
|
||||
nn.Linear(cond_dim, cond_channels),
|
||||
Rearrange("batch t -> batch t 1"),
|
||||
)
|
||||
|
||||
# make sure dimensions compatible
|
||||
self.residual_conv = (
|
||||
nn.Conv1d(in_channels, out_channels, 1)
|
||||
if in_channels != out_channels
|
||||
else nn.Identity()
|
||||
)
|
||||
|
||||
def forward(self, x, cond):
|
||||
"""
|
||||
x : [ batch_size x in_channels x horizon_steps ]
|
||||
cond : [ batch_size x cond_dim]
|
||||
|
||||
returns:
|
||||
out : [ batch_size x out_channels x horizon_steps ]
|
||||
"""
|
||||
out = self.blocks[0](x)
|
||||
embed = self.cond_encoder(cond)
|
||||
if self.cond_predict_scale:
|
||||
embed = embed.reshape(embed.shape[0], 2, self.out_channels, 1)
|
||||
scale = embed[:, 0, ...]
|
||||
bias = embed[:, 1, ...]
|
||||
out = scale * out + bias
|
||||
else:
|
||||
out = out + embed
|
||||
out = self.blocks[1](out)
|
||||
return out + self.residual_conv(x)
|
||||
|
||||
|
||||
class Unet1D(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transition_dim,
|
||||
cond_dim=None,
|
||||
diffusion_step_embed_dim=32,
|
||||
dim=32,
|
||||
dim_mults=(1, 2, 4, 8),
|
||||
smaller_encoder=False,
|
||||
cond_mlp_dims=None,
|
||||
kernel_size=5,
|
||||
n_groups=None,
|
||||
activation_type="Mish",
|
||||
cond_predict_scale=False,
|
||||
groupnorm_eps=1e-5,
|
||||
):
|
||||
super().__init__()
|
||||
dims = [transition_dim, *map(lambda m: dim * m, dim_mults)]
|
||||
in_out = list(zip(dims[:-1], dims[1:]))
|
||||
log.info(f"Channel dimensions: {in_out}")
|
||||
|
||||
dsed = diffusion_step_embed_dim
|
||||
self.time_mlp = nn.Sequential(
|
||||
SinusoidalPosEmb(dsed),
|
||||
nn.Linear(dsed, dsed * 4),
|
||||
nn.Mish(),
|
||||
nn.Linear(dsed * 4, dsed),
|
||||
)
|
||||
if cond_mlp_dims is not None:
|
||||
self.cond_mlp = ResidualMLP(
|
||||
dim_list=[cond_dim] + cond_mlp_dims,
|
||||
activation_type=activation_type,
|
||||
out_activation_type="Identity",
|
||||
)
|
||||
cond_block_dim = dsed + cond_mlp_dims[-1]
|
||||
else:
|
||||
cond_block_dim = dsed + cond_dim
|
||||
use_large_encoder_in_block = cond_mlp_dims is None and not smaller_encoder
|
||||
|
||||
mid_dim = dims[-1]
|
||||
self.mid_modules = nn.ModuleList(
|
||||
[
|
||||
ResidualBlock1D(
|
||||
mid_dim,
|
||||
mid_dim,
|
||||
cond_dim=cond_block_dim,
|
||||
kernel_size=kernel_size,
|
||||
n_groups=n_groups,
|
||||
cond_predict_scale=cond_predict_scale,
|
||||
larger_encoder=use_large_encoder_in_block,
|
||||
activation_type=activation_type,
|
||||
groupnorm_eps=groupnorm_eps,
|
||||
),
|
||||
ResidualBlock1D(
|
||||
mid_dim,
|
||||
mid_dim,
|
||||
cond_dim=cond_block_dim,
|
||||
kernel_size=kernel_size,
|
||||
n_groups=n_groups,
|
||||
cond_predict_scale=cond_predict_scale,
|
||||
larger_encoder=use_large_encoder_in_block,
|
||||
activation_type=activation_type,
|
||||
groupnorm_eps=groupnorm_eps,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
self.down_modules = nn.ModuleList([])
|
||||
for ind, (dim_in, dim_out) in enumerate(in_out):
|
||||
is_last = ind >= (len(in_out) - 1)
|
||||
self.down_modules.append(
|
||||
nn.ModuleList(
|
||||
[
|
||||
ResidualBlock1D(
|
||||
dim_in,
|
||||
dim_out,
|
||||
cond_dim=cond_block_dim,
|
||||
kernel_size=kernel_size,
|
||||
n_groups=n_groups,
|
||||
cond_predict_scale=cond_predict_scale,
|
||||
larger_encoder=use_large_encoder_in_block,
|
||||
activation_type=activation_type,
|
||||
groupnorm_eps=groupnorm_eps,
|
||||
),
|
||||
ResidualBlock1D(
|
||||
dim_out,
|
||||
dim_out,
|
||||
cond_dim=cond_block_dim,
|
||||
kernel_size=kernel_size,
|
||||
n_groups=n_groups,
|
||||
cond_predict_scale=cond_predict_scale,
|
||||
larger_encoder=use_large_encoder_in_block,
|
||||
activation_type=activation_type,
|
||||
groupnorm_eps=groupnorm_eps,
|
||||
),
|
||||
Downsample1d(dim_out) if not is_last else nn.Identity(),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
self.up_modules = nn.ModuleList([])
|
||||
for ind, (dim_in, dim_out) in enumerate(reversed(in_out[1:])):
|
||||
is_last = ind >= (len(in_out) - 1)
|
||||
self.up_modules.append(
|
||||
nn.ModuleList(
|
||||
[
|
||||
ResidualBlock1D(
|
||||
dim_out * 2,
|
||||
dim_in,
|
||||
cond_dim=cond_block_dim,
|
||||
kernel_size=kernel_size,
|
||||
n_groups=n_groups,
|
||||
cond_predict_scale=cond_predict_scale,
|
||||
larger_encoder=use_large_encoder_in_block,
|
||||
activation_type=activation_type,
|
||||
groupnorm_eps=groupnorm_eps,
|
||||
),
|
||||
ResidualBlock1D(
|
||||
dim_in,
|
||||
dim_in,
|
||||
cond_dim=cond_block_dim,
|
||||
kernel_size=kernel_size,
|
||||
n_groups=n_groups,
|
||||
cond_predict_scale=cond_predict_scale,
|
||||
larger_encoder=use_large_encoder_in_block,
|
||||
activation_type=activation_type,
|
||||
groupnorm_eps=groupnorm_eps,
|
||||
),
|
||||
Upsample1d(dim_in) if not is_last else nn.Identity(),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
self.final_conv = nn.Sequential(
|
||||
Conv1dBlock(
|
||||
dim,
|
||||
dim,
|
||||
kernel_size=kernel_size,
|
||||
n_groups=n_groups,
|
||||
activation_type=activation_type,
|
||||
eps=groupnorm_eps,
|
||||
),
|
||||
nn.Conv1d(dim, transition_dim, 1),
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
time,
|
||||
cond,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
x: (B,T,input_dim)
|
||||
time: (B,) or int, diffusion step
|
||||
cond: (B,obs_step,cond_dim)
|
||||
output: (B,T,input_dim)
|
||||
"""
|
||||
x = einops.rearrange(x, "b h t -> b t h")
|
||||
cond = cond.view(cond.shape[0], -1)
|
||||
if hasattr(self, "cond_mlp"):
|
||||
cond = self.cond_mlp(cond)
|
||||
|
||||
# 1. time
|
||||
if not torch.is_tensor(time):
|
||||
time = torch.tensor([time], dtype=torch.long, device=x.device)
|
||||
elif torch.is_tensor(time) and len(time.shape) == 0:
|
||||
time = time[None].to(x.device)
|
||||
# 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)
|
||||
|
||||
# encode local features
|
||||
h_local = list()
|
||||
h = []
|
||||
for idx, (resnet, resnet2, downsample) in enumerate(self.down_modules):
|
||||
x = resnet(x, global_feature)
|
||||
if idx == 0 and len(h_local) > 0:
|
||||
x = x + h_local[0]
|
||||
x = resnet2(x, global_feature)
|
||||
h.append(x)
|
||||
x = downsample(x)
|
||||
|
||||
for mid_module in self.mid_modules:
|
||||
x = mid_module(x, global_feature)
|
||||
|
||||
for idx, (resnet, resnet2, upsample) in enumerate(self.up_modules):
|
||||
x = torch.cat((x, h.pop()), dim=1)
|
||||
x = resnet(x, global_feature)
|
||||
if idx == len(self.up_modules) and len(h_local) > 0:
|
||||
x = x + h_local[1]
|
||||
x = resnet2(x, global_feature)
|
||||
x = upsample(x)
|
||||
|
||||
x = self.final_conv(x)
|
||||
|
||||
x = einops.rearrange(x, "b t h -> b h t")
|
||||
return x
|
||||
Reference in New Issue
Block a user