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
View File
+164
View File
@@ -0,0 +1,164 @@
"""
Critic networks.
"""
import torch
import copy
import einops
from copy import deepcopy
from model.common.mlp import MLP, ResidualMLP
from model.common.modules import SpatialEmb, RandomShiftsAug
class CriticObs(torch.nn.Module):
"""State-only critic network."""
def __init__(
self,
obs_dim,
mlp_dims,
activation_type="Mish",
use_layernorm=False,
residual_style=False,
**kwargs,
):
super().__init__()
mlp_dims = [obs_dim] + mlp_dims + [1]
if residual_style:
self.Q1 = ResidualMLP(
mlp_dims,
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
)
else:
self.Q1 = MLP(
mlp_dims,
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
verbose=False,
)
def forward(self, x):
x = x.view(x.size(0), -1)
q1 = self.Q1(x)
return q1
class CriticObsAct(torch.nn.Module):
"""State-action double critic network."""
def __init__(
self,
obs_dim,
mlp_dims,
action_dim,
action_steps=1,
activation_type="Mish",
use_layernorm=False,
residual_tyle=False,
**kwargs,
):
super().__init__()
mlp_dims = [obs_dim + action_dim * action_steps] + mlp_dims + [1]
if residual_tyle:
self.Q1 = ResidualMLP(
mlp_dims,
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
)
else:
self.Q1 = MLP(
mlp_dims,
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
verbose=False,
)
self.Q2 = copy.deepcopy(self.Q1)
def forward(self, x, action):
x = x.view(x.size(0), -1)
x = torch.cat((x, action), dim=-1)
q1 = self.Q1(x)
q2 = self.Q2(x)
return q1.squeeze(1), q2.squeeze(1)
class ViTCritic(CriticObs):
"""ViT + MLP, state only"""
def __init__(
self,
backbone,
obs_dim,
spatial_emb=128,
patch_repr_dim=128,
dropout=0,
augment=False,
num_img=1,
**kwargs,
):
# update input dim to mlp
mlp_obs_dim = spatial_emb * num_img + obs_dim
super().__init__(obs_dim=mlp_obs_dim, **kwargs)
self.backbone = backbone
if num_img > 1:
self.compress1 = SpatialEmb(
num_patch=121, # TODO: repr_dim // patch_repr_dim,
patch_dim=patch_repr_dim,
prop_dim=obs_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=obs_dim,
proj_dim=spatial_emb,
dropout=dropout,
)
if augment:
self.aug = RandomShiftsAug(pad=4)
self.augment = augment
def forward(
self,
obs: dict,
no_augment=False,
):
# flatten cond_dim if exists
if obs["rgb"].ndim == 5:
rgb = einops.rearrange(obs["rgb"], "b d c h w -> (b d) c h w")
else:
rgb = obs["rgb"]
if obs["state"].ndim == 3:
state = einops.rearrange(obs["state"], "b d c -> (b d) c")
else:
state = obs["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 and not no_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 and not no_augment:
rgb = self.aug(rgb) # uint8 -> float32
feat = self.backbone(rgb)
feat = self.compress.forward(feat, state)
feat = torch.cat([feat, state], dim=-1)
return super().forward(feat)
+97
View File
@@ -0,0 +1,97 @@
"""
Gaussian policy parameterization.
"""
import torch
import torch.distributions as D
import logging
log = logging.getLogger(__name__)
class GaussianModel(torch.nn.Module):
def __init__(
self,
network,
horizon_steps,
network_path=None,
device="cuda:0",
):
super().__init__()
self.device = device
self.network = network.to(device)
self.horizon_steps = horizon_steps
if network_path is not None:
checkpoint = torch.load(
network_path, map_location=self.device, weights_only=True
)
self.load_state_dict(
checkpoint["model"],
strict=False,
)
log.info("Loaded actor from %s", network_path)
log.info(
f"Number of network parameters: {sum(p.numel() for p in self.parameters())}"
)
def loss(self, true_action, cond, ent_coef):
B = len(true_action)
if isinstance(
cond, dict
): # image and state, only using one step observation right now
cond = cond[0]
else:
cond = cond[0].reshape(B, -1)
dist = self.forward_train(
cond,
deterministic=False,
)
true_action = true_action.view(B, -1)
loss = -dist.log_prob(true_action) # [B]
entropy = dist.entropy().mean()
loss = loss.mean() - entropy * ent_coef
return loss, {"entropy": entropy}
def forward_train(
self,
cond,
deterministic=False,
network_override=None,
):
"""
Calls the MLP to compute the mean, scale, and logits of the GMM. Returns the torch.Distribution object.
"""
if network_override is not None:
means, scales = network_override(cond)
else:
means, scales = self.network(cond)
if deterministic:
# low-noise for all Gaussian dists
scales = torch.ones_like(means) * 1e-4
return D.Normal(loc=means, scale=scales)
def forward(
self,
cond,
deterministic=False,
randn_clip_value=10,
network_override=None,
):
if isinstance(cond, dict):
B = cond["state"].shape[0]
else:
B = cond.shape[0]
T = self.horizon_steps
dist = self.forward_train(
cond,
deterministic=deterministic,
network_override=network_override,
)
sampled_action = dist.sample()
sampled_action.clamp_(
dist.loc - randn_clip_value * dist.scale,
dist.loc + randn_clip_value * dist.scale,
)
return sampled_action.view(B, T, -1)
+83
View File
@@ -0,0 +1,83 @@
"""
GMM policy parameterization.
"""
import torch
import torch.distributions as D
import logging
log = logging.getLogger(__name__)
class GMMModel(torch.nn.Module):
def __init__(
self,
network,
horizon_steps,
device="cuda:0",
**kwargs,
):
super().__init__()
self.device = device
self.network = network.to(device)
self.horizon_steps = horizon_steps
log.info(
f"Number of network parameters: {sum(p.numel() for p in self.parameters())}"
)
def loss(self, true_action, obs_cond, **kwargs):
B = len(true_action)
cond = obs_cond[0].reshape(B, -1)
dist, entropy, _ = self.forward_train(
cond,
deterministic=False,
)
true_action = true_action.view(B, -1)
loss = -dist.log_prob(true_action) # [B]
loss = loss.mean()
return loss, {"entropy": entropy}
def forward_train(
self,
cond,
deterministic=False,
):
"""
Calls the MLP to compute the mean, scale, and logits of the GMM. Returns the torch.Distribution object.
"""
means, scales, logits = self.network(cond)
if deterministic:
# low-noise for all Gaussian dists
scales = torch.ones_like(means) * 1e-4
# mixture components - make sure that `batch_shape` for the distribution is equal to (batch_size, num_modes) since MixtureSameFamily expects this shape
# Each mode has mean vector of dim T*D
component_distribution = D.Normal(loc=means, scale=scales)
component_distribution = D.Independent(component_distribution, 1)
component_entropy = component_distribution.entropy()
approx_entropy = torch.mean(
torch.sum(logits.softmax(-1) * component_entropy, dim=-1)
)
std = torch.mean(torch.sum(logits.softmax(-1) * scales.mean(-1), dim=-1))
# unnormalized logits to categorical distribution for mixing the modes
mixture_distribution = D.Categorical(logits=logits)
dist = D.MixtureSameFamily(
mixture_distribution=mixture_distribution,
component_distribution=component_distribution,
)
return dist, approx_entropy, std
def forward(self, cond, deterministic=False):
B = cond.shape[0]
T = self.horizon_steps
dist, _, _ = self.forward_train(
cond,
deterministic=deterministic,
)
sampled_action = dist.sample()
sampled_action = sampled_action.view(B, T, -1)
return sampled_action
+160
View File
@@ -0,0 +1,160 @@
"""
Implementation of Multi-layer Perception (MLP).
Residual model is taken from https://github.com/ALRhub/d3il/blob/main/agents/models/common/mlp.py
"""
import torch
from torch import nn
from torch.nn.utils import spectral_norm
from collections import OrderedDict
import logging
activation_dict = nn.ModuleDict(
{
"ReLU": nn.ReLU(),
"ELU": nn.ELU(),
"GELU": nn.GELU(),
"Tanh": nn.Tanh(),
"Mish": nn.Mish(),
"Identity": nn.Identity(),
"Softplus": nn.Softplus(),
}
)
class MLP(nn.Module):
def __init__(
self,
dim_list,
append_dim=0,
append_layers=None,
activation_type="Tanh",
out_activation_type="Identity",
use_layernorm=False,
use_spectralnorm=False,
verbose=False,
):
super(MLP, self).__init__()
# Construct module list: if use `Python List`, the modules are not
# added to computation graph. Instead, we should use `nn.ModuleList()`.
self.moduleList = nn.ModuleList()
self.append_layers = append_layers
num_layer = len(dim_list) - 1
for idx in range(num_layer):
i_dim = dim_list[idx]
o_dim = dim_list[idx + 1]
if append_dim > 0 and idx in append_layers:
i_dim += append_dim
linear_layer = nn.Linear(i_dim, o_dim)
if use_spectralnorm:
linear_layer = spectral_norm(linear_layer)
if idx == num_layer - 1:
module = nn.Sequential(
OrderedDict(
[
("linear_1", linear_layer),
("act_1", activation_dict[out_activation_type]),
]
)
)
else:
if use_layernorm:
module = nn.Sequential(
OrderedDict(
[
("linear_1", linear_layer),
("norm_1", nn.LayerNorm(o_dim)),
("act_1", activation_dict[activation_type]),
]
)
)
else:
module = nn.Sequential(
OrderedDict(
[
("linear_1", linear_layer),
("act_1", activation_dict[activation_type]),
]
)
)
self.moduleList.append(module)
if verbose:
logging.info(self.moduleList)
def forward(self, x, append=None):
for layer_ind, m in enumerate(self.moduleList):
if append is not None and layer_ind in self.append_layers:
x = torch.cat((x, append), dim=-1)
x = m(x)
return x
class ResidualMLP(nn.Module):
"""
Simple multi layer perceptron network with residual connections for
benchmarking the performance of different networks. The resiudal layers
are based on the IBC paper implementation, which uses 2 residual lalyers
with pre-actication with or without dropout and normalization.
"""
def __init__(
self,
dim_list,
activation_type="Mish",
out_activation_type="Identity",
use_layernorm=False,
):
super(ResidualMLP, self).__init__()
hidden_dim = dim_list[1]
num_hidden_layers = len(dim_list) - 3
assert num_hidden_layers % 2 == 0
self.layers = nn.ModuleList([nn.Linear(dim_list[0], hidden_dim)])
self.layers.extend(
[
TwoLayerPreActivationResNetLinear(
hidden_dim=hidden_dim,
activation_type=activation_type,
use_layernorm=use_layernorm,
)
for _ in range(1, num_hidden_layers, 2)
]
)
self.layers.append(nn.Linear(hidden_dim, dim_list[-1]))
self.layers.append(activation_dict[out_activation_type])
def forward(self, x):
for _, layer in enumerate(self.layers):
x = layer(x)
return x
class TwoLayerPreActivationResNetLinear(nn.Module):
def __init__(
self,
hidden_dim,
activation_type="Mish",
use_layernorm=False,
):
super().__init__()
self.l1 = nn.Linear(hidden_dim, hidden_dim)
self.l2 = nn.Linear(hidden_dim, hidden_dim)
self.act = activation_dict[activation_type]
if use_layernorm:
self.norm1 = nn.LayerNorm(hidden_dim, eps=1e-06)
self.norm2 = nn.LayerNorm(hidden_dim, eps=1e-06)
def forward(self, x):
x_input = x
if hasattr(self, "norm1"):
x = self.norm1(x)
x = self.l1(self.act(x))
if hasattr(self, "norm2"):
x = self.norm2(x)
x = self.l2(self.act(x))
return x + x_input
+248
View File
@@ -0,0 +1,248 @@
"""
MLP models for Gaussian policy.
"""
import torch
import torch.nn as nn
import einops
from copy import deepcopy
from model.common.mlp import MLP, ResidualMLP
from model.common.modules import SpatialEmb, RandomShiftsAug
class Gaussian_VisionMLP(nn.Module):
"""With ViT backbone"""
def __init__(
self,
backbone,
transition_dim,
horizon_steps,
cond_dim,
mlp_dims=[256, 256, 256],
activation_type="Mish",
residual_style=False,
use_layernorm=False,
fixed_std=None,
learn_fixed_std=False,
std_min=0.01,
std_max=1,
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(),
)
# head
self.transition_dim = transition_dim
self.horizon_steps = horizon_steps
input_dim = visual_feature_dim + cond_dim
output_dim = transition_dim * horizon_steps
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="Identity",
use_layernorm=use_layernorm,
)
if fixed_std is None:
self.mlp_logvar = MLP(
[input_dim] + mlp_dims[-1:] + [output_dim],
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
)
elif learn_fixed_std: # initialize to fixed_std
self.logvar = torch.nn.Parameter(
torch.log(torch.tensor([fixed_std**2 for _ in range(transition_dim)])),
requires_grad=True,
)
self.logvar_min = torch.nn.Parameter(
torch.log(torch.tensor(std_min**2)), requires_grad=False
)
self.logvar_max = torch.nn.Parameter(
torch.log(torch.tensor(std_max**2)), requires_grad=False
)
self.use_fixed_std = fixed_std is not None
self.fixed_std = fixed_std
self.learn_fixed_std = learn_fixed_std
def forward(self, x):
B = len(x["state"])
device = x["state"].device
# flatten cond_dim if exists
if x["rgb"].ndim == 5:
rgb = einops.rearrange(x["rgb"], "b d c h w -> (b d) c h w")
else:
rgb = x["rgb"]
if x["state"].ndim == 3:
state = einops.rearrange(x["state"], "b d c -> (b d) c")
else:
state = x["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)
# mlp
x_encoded = torch.cat([feat, state], dim=-1)
out_mean = self.mlp_mean(x_encoded)
out_mean = torch.tanh(out_mean).view(
B, self.horizon_steps * self.transition_dim
) # tanh squashing in [-1, 1]
if self.learn_fixed_std:
out_logvar = torch.clamp(self.logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
out_scale = out_scale.view(1, self.transition_dim)
out_scale = out_scale.repeat(B, self.horizon_steps)
elif self.use_fixed_std:
out_scale = torch.ones_like(out_mean).to(device) * self.fixed_std
else:
out_logvar = self.mlp_logvar(x_encoded).view(
B, self.horizon_steps * self.transition_dim
)
out_logvar = torch.clamp(out_logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
return out_mean, out_scale
class Gaussian_MLP(nn.Module):
def __init__(
self,
transition_dim,
horizon_steps,
cond_dim,
mlp_dims=[256, 256, 256],
activation_type="Mish",
residual_style=False,
use_layernorm=False,
fixed_std=None,
learn_fixed_std=False,
std_min=0.01,
std_max=1,
):
super().__init__()
self.transition_dim = transition_dim
self.horizon_steps = horizon_steps
input_dim = cond_dim
output_dim = transition_dim * horizon_steps
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="Identity",
use_layernorm=use_layernorm,
)
if fixed_std is None:
self.mlp_logvar = MLP(
[input_dim] + mlp_dims[-1:] + [output_dim],
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
)
elif learn_fixed_std: # initialize to fixed_std
self.logvar = torch.nn.Parameter(
torch.log(torch.tensor([fixed_std**2 for _ in range(transition_dim)])),
requires_grad=True,
)
self.logvar_min = torch.nn.Parameter(
torch.log(torch.tensor(std_min**2)), requires_grad=False
)
self.logvar_max = torch.nn.Parameter(
torch.log(torch.tensor(std_max**2)), requires_grad=False
)
self.use_fixed_std = fixed_std is not None
self.fixed_std = fixed_std
self.learn_fixed_std = learn_fixed_std
def forward(self, x):
B = len(x)
# mlp
out_mean = self.mlp_mean(x)
out_mean = torch.tanh(out_mean).view(
B, self.horizon_steps * self.transition_dim
) # tanh squashing in [-1, 1]
if self.learn_fixed_std:
out_logvar = torch.clamp(self.logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
out_scale = out_scale.view(1, self.transition_dim)
out_scale = out_scale.repeat(B, self.horizon_steps)
elif self.use_fixed_std:
out_scale = torch.ones_like(out_mean).to(x.device) * self.fixed_std
else:
out_logvar = self.mlp_logvar(x).view(
B, self.horizon_steps * self.transition_dim
)
out_logvar = torch.clamp(out_logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
return out_mean, out_scale
+106
View File
@@ -0,0 +1,106 @@
"""
MLP models for GMM policy.
"""
import torch
import torch.nn as nn
from model.common.mlp import MLP, ResidualMLP
class GMM_MLP(nn.Module):
def __init__(
self,
transition_dim,
horizon_steps,
cond_dim=None,
mlp_dims=[256, 256, 256],
num_modes=5,
activation_type="Mish",
residual_style=False,
use_layernorm=False,
fixed_std=None,
learn_fixed_std=False,
std_min=0.01,
std_max=1,
):
super().__init__()
self.transition_dim = transition_dim
self.horizon_steps = horizon_steps
input_dim = cond_dim
output_dim = transition_dim * horizon_steps * num_modes
self.num_modes = num_modes
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="Identity",
use_layernorm=use_layernorm,
)
if fixed_std is None:
self.mlp_logvar = model(
[input_dim] + mlp_dims + [output_dim],
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
)
elif (
learn_fixed_std
): # initialize to fixed_std, separate for each action and mode
self.logvar = torch.nn.Parameter(
torch.log(
torch.tensor(
[fixed_std**2 for _ in range(transition_dim * num_modes)]
)
),
requires_grad=True,
)
self.logvar_min = torch.nn.Parameter(
torch.log(torch.tensor(std_min**2)), requires_grad=False
)
self.logvar_max = torch.nn.Parameter(
torch.log(torch.tensor(std_max**2)), requires_grad=False
)
self.use_fixed_std = fixed_std is not None
self.fixed_std = fixed_std
self.learn_fixed_std = learn_fixed_std
# mode weights
self.mlp_weights = model(
[input_dim] + mlp_dims + [num_modes],
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
)
def forward(self, x):
B = len(x)
# mlp
out_mean = self.mlp_mean(x)
out_mean = torch.tanh(out_mean).view(
B, self.num_modes, self.horizon_steps * self.transition_dim
) # tanh squashing in [-1, 1]
if self.learn_fixed_std:
out_logvar = torch.clamp(self.logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
out_scale = out_scale.view(1, self.num_modes, self.transition_dim)
out_scale = out_scale.repeat(B, 1, self.horizon_steps)
elif self.use_fixed_std:
out_scale = torch.ones_like(out_mean).to(x.device) * self.fixed_std
else:
out_logvar = self.mlp_logvar(x).view(
B, self.num_modes, self.horizon_steps * self.transition_dim
)
out_logvar = torch.clamp(out_logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
out_weights = self.mlp_weights(x)
out_weights = out_weights.view(B, self.num_modes)
return out_mean, out_scale, out_weights
+69
View File
@@ -0,0 +1,69 @@
"""
Additional implementation of the ViT image encoder from https://github.com/hengyuan-hu/ibrl/tree/main
"""
import torch
import torch.nn as nn
class SpatialEmb(nn.Module):
def __init__(self, num_patch, patch_dim, prop_dim, proj_dim, dropout):
super().__init__()
proj_in_dim = num_patch + prop_dim
num_proj = patch_dim
self.patch_dim = patch_dim
self.prop_dim = prop_dim
self.input_proj = nn.Sequential(
nn.Linear(proj_in_dim, proj_dim),
nn.LayerNorm(proj_dim),
nn.ReLU(inplace=True),
)
self.weight = nn.Parameter(torch.zeros(1, num_proj, proj_dim))
self.dropout = nn.Dropout(dropout)
nn.init.normal_(self.weight)
def extra_repr(self) -> str:
return f"weight: nn.Parameter ({self.weight.size()})"
def forward(self, feat: torch.Tensor, prop: torch.Tensor):
feat = feat.transpose(1, 2)
if self.prop_dim > 0:
repeated_prop = prop.unsqueeze(1).repeat(1, feat.size(1), 1)
feat = torch.cat((feat, repeated_prop), dim=-1)
y = self.input_proj(feat)
z = (self.weight * y).sum(1)
z = self.dropout(z)
return z
class RandomShiftsAug:
def __init__(self, pad):
self.pad = pad
def __call__(self, x):
n, c, h, w = x.size()
assert h == w
padding = tuple([self.pad] * 4)
x = nn.functional.pad(x, padding, "replicate")
eps = 1.0 / (h + 2 * self.pad)
arange = torch.linspace(
-1.0 + eps, 1.0 - eps, h + 2 * self.pad, device=x.device, dtype=x.dtype
)[:h]
arange = arange.unsqueeze(0).repeat(h, 1).unsqueeze(2)
base_grid = torch.cat([arange, arange.transpose(1, 0)], dim=2)
base_grid = base_grid.unsqueeze(0).repeat(n, 1, 1, 1)
shift = torch.randint(
0, 2 * self.pad + 1, size=(n, 1, 1, 2), device=x.device, dtype=x.dtype
)
shift *= 2.0 / (h + 2 * self.pad)
grid = base_grid + shift
return nn.functional.grid_sample(
x, grid, padding_mode="zeros", align_corners=False
)
+415
View File
@@ -0,0 +1,415 @@
"""
Implementation of Transformer, parameterized as Gaussian and GMM.
Modified from https://github.com/real-stanford/diffusion_policy/blob/main/diffusion_policy/model/diffusion/transformer_for_diffusion.py
"""
import logging
import torch
import torch.nn as nn
from model.diffusion.modules import SinusoidalPosEmb
logger = logging.getLogger(__name__)
class Gaussian_Transformer(nn.Module):
def __init__(
self,
transition_dim,
horizon_steps,
cond_dim,
transformer_embed_dim=256,
transformer_num_heads=8,
transformer_num_layers=6,
transformer_activation="gelu",
p_drop_emb=0.0,
p_drop_attn=0.0,
fixed_std=None,
learn_fixed_std=False,
std_min=0.01,
std_max=1,
):
super().__init__()
self.transition_dim = transition_dim
self.horizon_steps = horizon_steps
output_dim = transition_dim
if fixed_std is None: # learn the logvar
output_dim *= 2 # mean and logvar
logger.info("Using learned std")
elif learn_fixed_std: # learn logvar
self.logvar = torch.nn.Parameter(
torch.log(torch.tensor([fixed_std**2 for _ in range(transition_dim)])),
requires_grad=True,
)
logger.info(f"Using fixed std {fixed_std} with learning")
else:
logger.info(f"Using fixed std {fixed_std} without learning")
self.logvar_min = torch.nn.Parameter(
torch.log(torch.tensor(std_min**2)), requires_grad=False
)
self.logvar_max = torch.nn.Parameter(
torch.log(torch.tensor(std_max**2)), requires_grad=False
)
self.learn_fixed_std = learn_fixed_std
self.fixed_std = fixed_std
self.transformer = Transformer(
output_dim,
horizon_steps,
cond_dim,
T_cond=1, # right now we assume only one step of observation everywhere
n_layer=transformer_num_layers,
n_head=transformer_num_heads,
n_emb=transformer_embed_dim,
p_drop_emb=p_drop_emb,
p_drop_attn=p_drop_attn,
activation=transformer_activation,
)
def forward(self, cond):
"""
cond: (B,cond_dim)
output: (B,horizon*transition)
"""
B = len(cond)
cond = cond.unsqueeze(1) # (B,1,cond_dim)
out, _ = self.transformer(cond) # (B,horizon,output_dim)
# use the first half of the output as mean
out_mean = torch.tanh(out[:, :, : self.transition_dim])
out_mean = out_mean.view(B, self.horizon_steps * self.transition_dim)
if self.learn_fixed_std:
out_logvar = torch.clamp(self.logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
out_scale = out_scale.view(1, self.transition_dim)
out_scale = out_scale.repeat(B, self.horizon_steps)
elif self.fixed_std is not None:
out_scale = torch.ones_like(out_mean).to(cond.device) * self.fixed_std
else:
out_logvar = out[:, :, self.transition_dim :]
out_logvar = out_logvar.reshape(B, self.horizon_steps * self.transition_dim)
out_logvar = torch.clamp(out_logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
return out_mean, out_scale
class GMM_Transformer(nn.Module):
def __init__(
self,
transition_dim,
horizon_steps,
cond_dim,
num_modes=5,
transformer_embed_dim=256,
transformer_num_heads=8,
transformer_num_layers=6,
transformer_activation="gelu",
p_drop_emb=0,
p_drop_attn=0,
fixed_std=None,
learn_fixed_std=False,
std_min=0.01,
std_max=1,
):
super().__init__()
self.num_modes = num_modes
self.transition_dim = transition_dim
self.horizon_steps = horizon_steps
output_dim = transition_dim * num_modes
# + num_modes # mean and modes
if fixed_std is None:
output_dim += num_modes * transition_dim # logvar for each mode
logger.info("Using learned std")
elif (
learn_fixed_std
): # initialize to fixed_std, separate for each action and mode, but same along horizon
self.logvar = torch.nn.Parameter(
torch.log(
torch.tensor(
[fixed_std**2 for _ in range(num_modes * transition_dim)]
)
),
requires_grad=True,
)
logger.info(f"Using fixed std {fixed_std} with learning")
else:
logger.info(f"Using fixed std {fixed_std} without learning")
self.logvar_min = torch.nn.Parameter(
torch.log(torch.tensor(std_min**2)), requires_grad=False
)
self.logvar_max = torch.nn.Parameter(
torch.log(torch.tensor(std_max**2)), requires_grad=False
)
self.fixed_std = fixed_std
self.learn_fixed_std = learn_fixed_std
self.transformer = Transformer(
output_dim,
horizon_steps,
cond_dim,
T_cond=1, # right now we assume only one step of observation everywhere
n_layer=transformer_num_layers,
n_head=transformer_num_heads,
n_emb=transformer_embed_dim,
p_drop_emb=p_drop_emb,
p_drop_attn=p_drop_attn,
activation=transformer_activation,
)
self.modes_head = nn.Linear(horizon_steps * transformer_embed_dim, num_modes)
def forward(self, cond):
"""
cond: (B,cond_dim)
output: (B,horizon*transition)
"""
B = len(cond)
cond = cond.unsqueeze(1) # (B,1,cond_dim)
out, out_prehead = self.transformer(
cond
) # (B,horizon,output_dim), (B,horizon,emb_dim)
# use the first half of the output as mean
out_mean = torch.tanh(out[:, :, : self.num_modes * self.transition_dim])
out_mean = out_mean.reshape(
B, self.horizon_steps, self.num_modes, self.transition_dim
)
out_mean = out_mean.permute(0, 2, 1, 3) # flip horizons and modes
out_mean = out_mean.reshape(
B, self.num_modes, self.horizon_steps * self.transition_dim
)
if self.learn_fixed_std:
out_logvar = torch.clamp(self.logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
out_scale = out_scale.view(1, self.num_modes, self.transition_dim)
out_scale = out_scale.repeat(B, 1, self.horizon_steps)
elif self.fixed_std is not None:
out_scale = torch.ones_like(out_mean).to(cond.device) * self.fixed_std
else:
out_logvar = out[
:, :, self.num_modes * self.transition_dim : -self.num_modes
]
out_logvar = out_logvar.reshape(
B, self.horizon_steps, self.num_modes, self.transition_dim
)
out_logvar = out_logvar.permute(0, 2, 1, 3) # flip horizons and modes
out_logvar = out_logvar.reshape(
B, self.num_modes, self.horizon_steps * self.transition_dim
)
out_logvar = torch.clamp(out_logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
# use last horizon step as the mode weights - as it depends on the entire context
# out_weights = out[:, -1, -self.num_modes :] # (B,num_modes)
out_weights = self.modes_head(out_prehead.view(B, -1))
return out_mean, out_scale, out_weights
class Transformer(nn.Module):
def __init__(
self,
output_dim,
horizon,
cond_dim,
T_cond=1,
n_layer=12,
n_head=12,
n_emb=768,
p_drop_emb=0.0,
p_drop_attn=0.0,
causal_attn=False,
n_cond_layers=0,
activation="gelu",
):
super().__init__()
# encoder for observations
self.cond_obs_emb = nn.Linear(cond_dim, n_emb)
self.cond_pos_emb = nn.Parameter(torch.zeros(1, T_cond, n_emb))
if n_cond_layers > 0:
encoder_layer = nn.TransformerEncoderLayer(
d_model=n_emb,
nhead=n_head,
dim_feedforward=4 * n_emb,
dropout=p_drop_attn,
activation=activation,
batch_first=True,
norm_first=True,
)
self.encoder = nn.TransformerEncoder(
encoder_layer=encoder_layer,
num_layers=n_cond_layers,
)
else:
self.encoder = nn.Sequential(
nn.Linear(n_emb, 4 * n_emb),
nn.Mish(),
nn.Linear(4 * n_emb, n_emb),
)
# decoder
self.pos_emb = nn.Parameter(torch.zeros(1, horizon, n_emb))
self.drop = nn.Dropout(p_drop_emb)
decoder_layer = nn.TransformerDecoderLayer(
d_model=n_emb,
nhead=n_head,
dim_feedforward=4 * n_emb,
dropout=p_drop_attn,
activation=activation,
batch_first=True,
norm_first=True, # important for stability
)
self.decoder = nn.TransformerDecoder(
decoder_layer=decoder_layer, num_layers=n_layer
)
# attention mask
if causal_attn:
# causal mask to ensure that attention is only applied to the left in the input sequence
# torch.nn.Transformer uses additive mask as opposed to multiplicative mask in minGPT
# therefore, the upper triangle should be -inf and others (including diag) should be 0.
sz = horizon
mask = (torch.triu(torch.ones(sz, sz)) == 1).transpose(0, 1)
mask = (
mask.float()
.masked_fill(mask == 0, float("-inf"))
.masked_fill(mask == 1, float(0.0))
)
self.register_buffer("mask", mask)
t, s = torch.meshgrid(
torch.arange(horizon), torch.arange(T_cond), indexing="ij"
)
mask = t >= (
s - 1
) # add one dimension since time is the first token in cond
mask = (
mask.float()
.masked_fill(mask == 0, float("-inf"))
.masked_fill(mask == 1, float(0.0))
)
self.register_buffer("memory_mask", mask)
else:
self.mask = None
self.memory_mask = None
# decoder head
self.ln_f = nn.LayerNorm(n_emb)
self.head = nn.Linear(n_emb, output_dim)
# constants
self.T_cond = T_cond
self.horizon = horizon
# init
self.apply(self._init_weights)
def _init_weights(self, module):
ignore_types = (
nn.Dropout,
SinusoidalPosEmb,
nn.TransformerEncoderLayer,
nn.TransformerDecoderLayer,
nn.TransformerEncoder,
nn.TransformerDecoder,
nn.ModuleList,
nn.Mish,
nn.Sequential,
)
if isinstance(module, (nn.Linear, nn.Embedding)):
torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
torch.nn.init.zeros_(module.bias)
elif isinstance(module, nn.MultiheadAttention):
weight_names = [
"in_proj_weight",
"q_proj_weight",
"k_proj_weight",
"v_proj_weight",
]
for name in weight_names:
weight = getattr(module, name)
if weight is not None:
torch.nn.init.normal_(weight, mean=0.0, std=0.02)
bias_names = ["in_proj_bias", "bias_k", "bias_v"]
for name in bias_names:
bias = getattr(module, name)
if bias is not None:
torch.nn.init.zeros_(bias)
elif isinstance(module, nn.LayerNorm):
torch.nn.init.zeros_(module.bias)
torch.nn.init.ones_(module.weight)
elif isinstance(module, Transformer):
torch.nn.init.normal_(module.pos_emb, mean=0.0, std=0.02)
if module.cond_obs_emb is not None:
torch.nn.init.normal_(module.cond_pos_emb, mean=0.0, std=0.02)
elif isinstance(module, ignore_types):
# no param
pass
else:
raise RuntimeError("Unaccounted module {}".format(module))
def forward(
self,
cond: torch.Tensor,
**kwargs,
):
"""
cond: (B, T, cond_dim)
output: (B, T, output_dim)
"""
# encoder
cond_embeddings = self.cond_obs_emb(cond) # (B,To,n_emb)
tc = cond_embeddings.shape[1]
position_embeddings = self.cond_pos_emb[
:, :tc, :
] # each position maps to a (learnable) vector
x = self.drop(cond_embeddings + position_embeddings)
x = self.encoder(x)
memory = x
# (B,T_cond,n_emb)
# decoder
position_embeddings = self.pos_emb[
:, : self.horizon, :
] # each position maps to a (learnable) vector
position_embeddings = position_embeddings.expand(
cond.shape[0], self.horizon, -1
) # repeat for batch dimension
x = self.drop(position_embeddings)
# (B,T,n_emb)
x = self.decoder(
tgt=x,
memory=memory,
tgt_mask=self.mask,
memory_mask=self.memory_mask,
)
# (B,T,n_emb)
# head
x_prehead = self.ln_f(x)
x = self.head(x_prehead)
# (B,T,n_out)
return x, x_prehead
if __name__ == "__main__":
transformer = Transformer(
output_dim=10,
horizon=4,
T_cond=1,
cond_dim=16,
causal_attn=False, # no need to use for delta control
# From Cheng: I found the causal attention masking to be critical to get the transformer variant of diffusion policy to work. My suspicion is that when used without it, the model "cheats" by looking ahead into future end-effector poses, which is almost identical to the action of the current timestep.
n_cond_layers=0,
)
# opt = transformer.configure_optimizers()
cond = torch.zeros((4, 1, 16)) # B x 1 x cond_dim
out, _ = transformer(cond)
+236
View File
@@ -0,0 +1,236 @@
"""
ViT image encoder implementation from IBRL, https://github.com/hengyuan-hu/ibrl
"""
from dataclasses import dataclass, field
from typing import List
import einops
import torch
from torch import nn
from torch.nn.init import trunc_normal_
@dataclass
class VitEncoderConfig:
patch_size: int = 8
depth: int = 1
embed_dim: int = 128
num_heads: int = 4
act_layer = nn.GELU
stride: int = -1
embed_style: str = "embed2"
embed_norm: int = 0
class VitEncoder(nn.Module):
def __init__(self, obs_shape: List[int], cfg: VitEncoderConfig):
super().__init__()
self.obs_shape = obs_shape
self.cfg = cfg
self.vit = MinVit(
embed_style=cfg.embed_style,
embed_dim=cfg.embed_dim,
embed_norm=cfg.embed_norm,
num_head=cfg.num_heads,
depth=cfg.depth,
)
self.num_patch = self.vit.num_patches
self.patch_repr_dim = self.cfg.embed_dim
self.repr_dim = self.cfg.embed_dim * self.vit.num_patches
def forward(self, obs, flatten=False) -> torch.Tensor:
# assert obs.max() > 5
obs = obs / 255.0 - 0.5
feats: torch.Tensor = self.vit.forward(obs)
if flatten:
feats = feats.flatten(1, 2)
return feats
class PatchEmbed1(nn.Module):
def __init__(self, embed_dim):
super().__init__()
self.conv = nn.Conv2d(3, embed_dim, kernel_size=8, stride=8)
self.num_patch = 144
self.patch_dim = embed_dim
def forward(self, x: torch.Tensor):
y = self.conv(x)
y = einops.rearrange(y, "b c h w -> b (h w) c")
return y
class PatchEmbed2(nn.Module):
def __init__(self, embed_dim, use_norm):
super().__init__()
layers = [
nn.Conv2d(3, embed_dim, kernel_size=8, stride=4),
nn.GroupNorm(embed_dim, embed_dim) if use_norm else nn.Identity(),
nn.ReLU(),
nn.Conv2d(embed_dim, embed_dim, kernel_size=3, stride=2),
]
self.embed = nn.Sequential(*layers)
self.num_patch = 121 # TODO: specifically for 96x96 set by Hengyuan?
self.patch_dim = embed_dim
def forward(self, x: torch.Tensor):
y = self.embed(x)
y = einops.rearrange(y, "b c h w -> b (h w) c")
return y
class MultiHeadAttention(nn.Module):
def __init__(self, embed_dim, num_head):
super().__init__()
assert embed_dim % num_head == 0
self.num_head = num_head
self.qkv_proj = nn.Linear(embed_dim, 3 * embed_dim)
self.out_proj = nn.Linear(embed_dim, embed_dim)
def forward(self, x, attn_mask):
"""
x: [batch, seq, embed_dim]
"""
qkv = self.qkv_proj(x)
q, k, v = einops.rearrange(
qkv, "b t (k h d) -> b k h t d", k=3, h=self.num_head
).unbind(1)
# force flash/mem-eff attention, it will raise error if flash cannot be applied
with torch.backends.cuda.sdp_kernel(enable_math=False):
attn_v = torch.nn.functional.scaled_dot_product_attention(
q, k, v, dropout_p=0.0, attn_mask=attn_mask
)
attn_v = einops.rearrange(attn_v, "b h t d -> b t (h d)")
return self.out_proj(attn_v)
class TransformerLayer(nn.Module):
def __init__(self, embed_dim, num_head, dropout):
super().__init__()
self.layer_norm1 = nn.LayerNorm(embed_dim)
self.mha = MultiHeadAttention(embed_dim, num_head)
self.layer_norm2 = nn.LayerNorm(embed_dim)
self.linear1 = nn.Linear(embed_dim, 4 * embed_dim)
self.linear2 = nn.Linear(4 * embed_dim, embed_dim)
self.dropout = nn.Dropout(dropout)
def forward(self, x, attn_mask=None):
x = x + self.dropout(self.mha(self.layer_norm1(x), attn_mask))
x = x + self.dropout(self._ff_block(self.layer_norm2(x)))
return x
def _ff_block(self, x):
x = self.linear2(nn.functional.gelu(self.linear1(x)))
return x
class MinVit(nn.Module):
def __init__(self, embed_style, embed_dim, embed_norm, num_head, depth):
super().__init__()
if embed_style == "embed1":
self.patch_embed = PatchEmbed1(embed_dim)
elif embed_style == "embed2":
self.patch_embed = PatchEmbed2(embed_dim, use_norm=embed_norm)
else:
assert False
self.pos_embed = nn.Parameter(
torch.zeros(1, self.patch_embed.num_patch, embed_dim)
)
layers = [
TransformerLayer(embed_dim, num_head, dropout=0) for _ in range(depth)
]
self.net = nn.Sequential(*layers)
self.norm = nn.LayerNorm(embed_dim)
self.num_patches = self.patch_embed.num_patch
# weight init
trunc_normal_(self.pos_embed, std=0.02)
named_apply(init_weights_vit_timm, self)
def forward(self, x):
x = self.patch_embed(x)
x = x + self.pos_embed
x = self.net(x)
return self.norm(x)
def init_weights_vit_timm(module: nn.Module, name: str = ""):
"""ViT weight initialization, original timm impl (for reproducibility)"""
if isinstance(module, nn.Linear):
trunc_normal_(module.weight, std=0.02)
if module.bias is not None:
nn.init.zeros_(module.bias)
def named_apply(
fn, module: nn.Module, name="", depth_first=True, include_root=False
) -> nn.Module:
if not depth_first and include_root:
fn(module=module, name=name)
for child_name, child_module in module.named_children():
child_name = ".".join((name, child_name)) if name else child_name
named_apply(
fn=fn,
module=child_module,
name=child_name,
depth_first=depth_first,
include_root=True,
)
if depth_first and include_root:
fn(module=module, name=name)
return module
def test_patch_embed():
print("embed 1")
embed = PatchEmbed1(128)
x = torch.rand(10, 3, 96, 96)
y = embed(x)
print(y.size())
print("embed 2")
embed = PatchEmbed2(128, True)
x = torch.rand(10, 3, 96, 96)
y = embed(x)
print(y.size())
def test_transformer_layer():
embed = PatchEmbed1(128)
x = torch.rand(10, 3, 96, 96)
y = embed(x)
print(y.size())
transformer = TransformerLayer(128, 4, False, 0)
z = transformer(y)
print(z.size())
if __name__ == "__main__":
import rich.traceback
import pyrallis
@dataclass
class MainConfig:
net_type: str = "vit"
obs_shape: list[int] = field(default_factory=lambda: [3, 96, 96])
vit: VitEncoderConfig = field(default_factory=lambda: VitEncoderConfig())
rich.traceback.install()
cfg = pyrallis.parse(config_class=MainConfig) # type: ignore
enc = VitEncoder(cfg.obs_shape, cfg.vit)
print(enc)
x = torch.rand(1, *cfg.obs_shape) * 255
print("output size:", enc(x, flatten=False).size())
print("repr dim:", enc.repr_dim, ", real dim:", enc(x, flatten=True).size())
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
View File
+30
View File
@@ -0,0 +1,30 @@
"""
Advantage-weighted regression (AWR) for Gaussian policy.
"""
import torch
import logging
from model.rl.gaussian_rwr import RWR_Gaussian
log = logging.getLogger(__name__)
class AWR_Gaussian(RWR_Gaussian):
def __init__(
self,
actor,
critic,
**kwargs,
):
super().__init__(actor=actor, **kwargs)
self.critic = critic.to(self.device)
def loss_critic(self, obs, advantages):
# get advantage
adv = self.critic(obs)
# Update critic
loss_critic = torch.mean((adv - advantages) ** 2)
return loss_critic
+120
View File
@@ -0,0 +1,120 @@
"""
PPO for Gaussian policy.
"""
from typing import Optional
import torch
from model.rl.gaussian_vpg import VPG_Gaussian
class PPO_Gaussian(VPG_Gaussian):
def __init__(
self,
clip_ploss_coef: float,
clip_vloss_coef: Optional[float] = None,
norm_adv: Optional[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
# Clipping value for value loss
self.clip_vloss_coef = clip_vloss_coef
def loss(
self,
obs,
actions,
returns,
oldvalues,
advantages,
oldlogprobs,
use_bc_loss=False,
):
"""
PPO loss
obs: (B, obs_step, obs_dim)
actions: (B, horizon_step, action_dim)
returns: (B, )
values: (B, )
advantages: (B,)
oldlogprobs: (B, )
"""
newlogprobs, entropy, std = self.get_logprobs(obs, actions)
newlogprobs = newlogprobs.clamp(min=-5, max=2)
oldlogprobs = oldlogprobs.clamp(min=-5, max=2)
entropy_loss = -entropy
# get ratio
logratio = newlogprobs - oldlogprobs
ratio = logratio.exp()
# get kl difference and whether value clipped
with torch.no_grad():
approx_kl = ((ratio - 1) - logratio).nanmean()
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()
bc_loss = 0.0
if use_bc_loss:
# See Eqn. 2 of https://arxiv.org/pdf/2403.03949.pdf
# Give a reward for maximizing probability of teacher policy's action with current policy.
# Actions are chosen along trajectory induced by current policy.
# Get counterfactual teacher actions
samples = self.forward(
cond=obs.float()
.unsqueeze(1)
.to(self.device), # B x horizon=1 x obs_dim
deterministic=False,
use_base_policy=True,
)
# Get logprobs of teacher actions under this policy
bc_logprobs, _, _ = self.get_logprobs(obs, samples, use_base_policy=False)
bc_logprobs = bc_logprobs.clamp(min=-5, max=2)
bc_loss = -bc_logprobs.mean()
return (
pg_loss,
entropy_loss,
v_loss,
clipfrac,
approx_kl.item(),
ratio.mean().item(),
bc_loss,
std.item(),
)
+58
View File
@@ -0,0 +1,58 @@
"""
Reward-weighted regression (RWR) for Gaussian policy.
"""
import torch
import logging
from model.common.gaussian import GaussianModel
import torch.distributions as D
log = logging.getLogger(__name__)
class RWR_Gaussian(GaussianModel):
def __init__(
self,
actor,
randn_clip_value=10,
**kwargs,
):
super().__init__(network=actor, **kwargs)
# assign actor
self.actor = self.network
# 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
# override
def loss(self, actions, obs, reward_weights):
cond = obs
B = cond.shape[0]
means, scales = self.network(cond)
dist = D.Normal(loc=means, scale=scales)
log_prob = dist.log_prob(actions.view(B, -1)).mean(-1)
log_prob = log_prob * reward_weights
log_prob = -log_prob.mean()
return log_prob
# override
@torch.no_grad()
def forward(self, cond, deterministic=False, **kwargs):
"""
Args:
cond: (batch_size, horizon, obs_dim)
Return:
actions: (batch_size, horizon_steps, transition_dim)
"""
B = cond.shape[0]
actions = super().forward(
cond=cond.view(B, -1),
deterministic=deterministic,
randn_clip_value=self.randn_clip_value,
)
return actions
+87
View File
@@ -0,0 +1,87 @@
"""
Policy gradient for Gaussian policy
"""
import torch
from copy import deepcopy
import logging
from model.common.gaussian import GaussianModel
class VPG_Gaussian(GaussianModel):
def __init__(
self,
actor,
critic,
cond_steps=1,
randn_clip_value=10,
network_path=None,
**kwargs,
):
super().__init__(network=actor, **kwargs)
self.cond_steps = cond_steps
self.randn_clip_value = randn_clip_value
# Value function for obs - simple MLP
self.critic = critic.to(self.device)
if network_path is not None:
checkpoint = torch.load(
network_path, map_location=self.device, weights_only=True
)
self.load_state_dict(
checkpoint["model"],
strict=False,
)
logging.info("Loaded actor from %s", network_path)
# Re-name network to actor
self.actor_ft = actor
# Save a copy of original actor
self.actor = deepcopy(actor)
for param in self.actor.parameters():
param.requires_grad = False
def get_logprobs(
self,
cond,
actions,
use_base_policy=False,
):
B, T, D = actions.shape
if not isinstance(cond, dict):
cond = cond.view(B, -1)
dist = self.forward_train(
cond,
deterministic=False,
network_override=self.actor if use_base_policy else None,
)
log_prob = dist.log_prob(actions.view(B, -1))
log_prob = log_prob.mean(-1)
entropy = dist.entropy().mean()
std = dist.scale.mean()
return log_prob, entropy, std
def loss(self, obs, actions, reward):
raise NotImplementedError
@torch.no_grad()
def forward(
self,
cond,
deterministic=False,
use_base_policy=False,
):
if isinstance(cond, dict):
B = cond["state"].shape[0]
else:
B = cond.shape[0]
cond = cond.view(B, -1)
return super().forward(
cond=cond,
deterministic=deterministic,
randn_clip_value=self.randn_clip_value,
network_override=self.actor if use_base_policy else None,
)
+102
View File
@@ -0,0 +1,102 @@
"""
PPO for GMM policy.
"""
from typing import Optional
import torch
from model.rl.gmm_vpg import VPG_GMM
class PPO_GMM(VPG_GMM):
def __init__(
self,
clip_ploss_coef: float,
clip_vloss_coef: Optional[float] = None,
norm_adv: Optional[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
# Clipping value for value loss
self.clip_vloss_coef = clip_vloss_coef
def loss(
self,
obs,
actions,
returns,
oldvalues,
advantages,
oldlogprobs,
**kwargs,
):
"""
PPO loss
obs: (B, obs_step, obs_dim)
actions: (B, horizon_step, action_dim)
returns: (B, )
values: (B, )
advantages: (B,)
oldlogprobs: (B, )
"""
newlogprobs, entropy, std = self.get_logprobs(obs, actions)
newlogprobs = newlogprobs.clamp(min=-5, max=2)
oldlogprobs = oldlogprobs.clamp(min=-5, max=2)
entropy_loss = -entropy.mean()
# get ratio
logratio = newlogprobs - oldlogprobs
ratio = logratio.exp()
# get kl difference and whether value clipped
with torch.no_grad():
approx_kl = ((ratio - 1) - logratio).nanmean()
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()
bc_loss = 0
return (
pg_loss,
entropy_loss,
v_loss,
clipfrac,
approx_kl.item(),
ratio.mean().item(),
bc_loss,
std.item(),
)
+56
View File
@@ -0,0 +1,56 @@
import torch
import logging
from model.common.gmm import GMMModel
class VPG_GMM(GMMModel):
def __init__(
self,
actor,
critic,
cond_steps=1,
network_path=None,
**kwargs,
):
super().__init__(network=actor, **kwargs)
self.cond_steps = cond_steps
# Re-name network to actor
self.actor_ft = actor
# Value function for obs - simple MLP
self.critic = critic.to(self.device)
if network_path is not None:
checkpoint = torch.load(
network_path, map_location=self.device, weights_only=True
)
self.load_state_dict(
checkpoint["model"],
strict=False,
)
logging.info("Loaded actor from %s", network_path)
def get_logprobs(
self,
cond,
actions,
):
B, T, D = actions.shape
dist, entropy, std = self.forward_train(
cond.view(B, -1),
deterministic=False,
)
log_prob = dist.log_prob(actions.view(B, -1))
return log_prob, entropy, std
def loss(self, obs, chains, reward):
raise NotImplementedError
# override to diffuse over action only
@torch.no_grad()
def forward(self, cond, deterministic=False):
B = cond.shape[0]
return super().forward(
cond=cond.view(B, -1),
deterministic=deterministic,
)