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