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