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())
|
||||
Reference in New Issue
Block a user