This commit is contained in:
allenzren
2024-09-03 21:03:27 -04:00
commit 8293b0936b
282 changed files with 34664 additions and 0 deletions
View File
+318
View File
@@ -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
)
+34
View File
@@ -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
+118
View File
@@ -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
+173
View File
@@ -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
+197
View File
@@ -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()
+193
View File
@@ -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(),
)
+139
View File
@@ -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,
)
+110
View File
@@ -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)
+153
View File
@@ -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
+465
View File
@@ -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
+164
View File
@@ -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)
+188
View File
@@ -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
+233
View File
@@ -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)
+95
View File
@@ -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)
+37
View File
@@ -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
+213
View File
@@ -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
+318
View File
@@ -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