squash commits
This commit is contained in:
+70
-27
@@ -3,6 +3,7 @@ Critic networks.
|
||||
|
||||
"""
|
||||
|
||||
from typing import Union
|
||||
import torch
|
||||
import copy
|
||||
import einops
|
||||
@@ -17,7 +18,7 @@ class CriticObs(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
obs_dim,
|
||||
cond_dim,
|
||||
mlp_dims,
|
||||
activation_type="Mish",
|
||||
use_layernorm=False,
|
||||
@@ -25,7 +26,7 @@ class CriticObs(torch.nn.Module):
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
mlp_dims = [obs_dim] + mlp_dims + [1]
|
||||
mlp_dims = [cond_dim] + mlp_dims + [1]
|
||||
if residual_style:
|
||||
self.Q1 = ResidualMLP(
|
||||
mlp_dims,
|
||||
@@ -42,9 +43,20 @@ class CriticObs(torch.nn.Module):
|
||||
verbose=False,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = x.view(x.size(0), -1)
|
||||
q1 = self.Q1(x)
|
||||
def forward(self, cond: Union[dict, torch.Tensor]):
|
||||
"""
|
||||
cond: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
or (B, num_feature) from ViT encoder
|
||||
"""
|
||||
if isinstance(cond, dict):
|
||||
B = len(cond["state"])
|
||||
|
||||
# flatten history
|
||||
state = cond["state"].view(B, -1)
|
||||
else:
|
||||
state = cond
|
||||
q1 = self.Q1(state)
|
||||
return q1
|
||||
|
||||
|
||||
@@ -53,7 +65,7 @@ class CriticObsAct(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
obs_dim,
|
||||
cond_dim,
|
||||
mlp_dims,
|
||||
action_dim,
|
||||
action_steps=1,
|
||||
@@ -63,7 +75,7 @@ class CriticObsAct(torch.nn.Module):
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
mlp_dims = [obs_dim + action_dim * action_steps] + mlp_dims + [1]
|
||||
mlp_dims = [cond_dim + action_dim * action_steps] + mlp_dims + [1]
|
||||
if residual_tyle:
|
||||
self.Q1 = ResidualMLP(
|
||||
mlp_dims,
|
||||
@@ -81,9 +93,21 @@ class CriticObsAct(torch.nn.Module):
|
||||
)
|
||||
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)
|
||||
def forward(self, cond: dict, action):
|
||||
"""
|
||||
cond: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
action: (B, Ta, Da)
|
||||
"""
|
||||
B = len(cond["state"])
|
||||
|
||||
# flatten history
|
||||
state = cond["state"].view(B, -1)
|
||||
|
||||
# flatten action
|
||||
action = action.view(B, -1)
|
||||
|
||||
x = torch.cat((state, action), dim=-1)
|
||||
q1 = self.Q1(x)
|
||||
q2 = self.Q2(x)
|
||||
return q1.squeeze(1), q2.squeeze(1)
|
||||
@@ -95,7 +119,8 @@ class ViTCritic(CriticObs):
|
||||
def __init__(
|
||||
self,
|
||||
backbone,
|
||||
obs_dim,
|
||||
cond_dim,
|
||||
img_cond_steps=1,
|
||||
spatial_emb=128,
|
||||
patch_repr_dim=128,
|
||||
dropout=0,
|
||||
@@ -104,14 +129,16 @@ class ViTCritic(CriticObs):
|
||||
**kwargs,
|
||||
):
|
||||
# update input dim to mlp
|
||||
mlp_obs_dim = spatial_emb * num_img + obs_dim
|
||||
super().__init__(obs_dim=mlp_obs_dim, **kwargs)
|
||||
mlp_obs_dim = spatial_emb * num_img + cond_dim
|
||||
super().__init__(cond_dim=mlp_obs_dim, **kwargs)
|
||||
self.backbone = backbone
|
||||
self.num_img = num_img
|
||||
self.img_cond_steps = img_cond_steps
|
||||
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,
|
||||
prop_dim=cond_dim,
|
||||
proj_dim=spatial_emb,
|
||||
dropout=dropout,
|
||||
)
|
||||
@@ -120,7 +147,7 @@ class ViTCritic(CriticObs):
|
||||
self.compress = SpatialEmb(
|
||||
num_patch=121,
|
||||
patch_dim=patch_repr_dim,
|
||||
prop_dim=obs_dim,
|
||||
prop_dim=cond_dim,
|
||||
proj_dim=spatial_emb,
|
||||
dropout=dropout,
|
||||
)
|
||||
@@ -130,23 +157,39 @@ class ViTCritic(CriticObs):
|
||||
|
||||
def forward(
|
||||
self,
|
||||
obs: dict,
|
||||
cond: 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")
|
||||
"""
|
||||
cond: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
rgb: (B, To, C, H, W)
|
||||
no_augment: whether to skip augmentation
|
||||
|
||||
TODO long term: more flexible handling of cond
|
||||
"""
|
||||
B, T_rgb, C, H, W = cond["rgb"].shape
|
||||
|
||||
# flatten history
|
||||
state = cond["state"].view(B, -1)
|
||||
|
||||
# Take recent images --- sometimes we want to use fewer img_cond_steps than cond_steps (e.g., 1 image but 3 prio)
|
||||
rgb = cond["rgb"][:, -self.img_cond_steps :]
|
||||
|
||||
# concatenate images in cond by channels
|
||||
if self.num_img > 1:
|
||||
rgb = rgb.reshape(B, T_rgb, self.num_img, 3, H, W)
|
||||
rgb = einops.rearrange(rgb, "b t n c h w -> b n (t 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"]
|
||||
rgb = einops.rearrange(rgb, "b t c h w -> b (t c) h w")
|
||||
|
||||
# convert rgb to float32 for augmentation
|
||||
rgb = rgb.float()
|
||||
|
||||
# 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.num_img > 1: # TODO: properly handle multiple images
|
||||
rgb1 = rgb[:, 0]
|
||||
rgb2 = rgb[:, 1]
|
||||
if self.augment and not no_augment:
|
||||
rgb1 = self.aug(rgb1)
|
||||
rgb2 = self.aug(rgb2)
|
||||
|
||||
@@ -22,7 +22,6 @@ class GaussianModel(torch.nn.Module):
|
||||
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
|
||||
@@ -35,15 +34,15 @@ class GaussianModel(torch.nn.Module):
|
||||
log.info(
|
||||
f"Number of network parameters: {sum(p.numel() for p in self.parameters())}"
|
||||
)
|
||||
self.horizon_steps = horizon_steps
|
||||
|
||||
def loss(self, true_action, cond, ent_coef):
|
||||
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,
|
||||
@@ -79,10 +78,7 @@ class GaussianModel(torch.nn.Module):
|
||||
randn_clip_value=10,
|
||||
network_override=None,
|
||||
):
|
||||
if isinstance(cond, dict):
|
||||
B = cond["state"].shape[0]
|
||||
else:
|
||||
B = cond.shape[0]
|
||||
B = len(cond["state"]) if "state" in cond else len(cond["rgb"])
|
||||
T = self.horizon_steps
|
||||
dist = self.forward_train(
|
||||
cond,
|
||||
|
||||
+18
-4
@@ -16,20 +16,34 @@ class GMMModel(torch.nn.Module):
|
||||
self,
|
||||
network,
|
||||
horizon_steps,
|
||||
network_path=None,
|
||||
device="cuda:0",
|
||||
**kwargs,
|
||||
):
|
||||
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,
|
||||
)
|
||||
logging.info("Loaded actor from %s", network_path)
|
||||
log.info(
|
||||
f"Number of network parameters: {sum(p.numel() for p in self.parameters())}"
|
||||
)
|
||||
self.horizon_steps = horizon_steps
|
||||
|
||||
def loss(self, true_action, obs_cond, **kwargs):
|
||||
def loss(
|
||||
self,
|
||||
true_action,
|
||||
cond,
|
||||
**kwargs,
|
||||
):
|
||||
B = len(true_action)
|
||||
cond = obs_cond[0].reshape(B, -1)
|
||||
dist, entropy, _ = self.forward_train(
|
||||
cond,
|
||||
deterministic=False,
|
||||
@@ -72,7 +86,7 @@ class GMMModel(torch.nn.Module):
|
||||
return dist, approx_entropy, std
|
||||
|
||||
def forward(self, cond, deterministic=False):
|
||||
B = cond.shape[0]
|
||||
B = len(cond["state"]) if "state" in cond else len(cond["rgb"])
|
||||
T = self.horizon_steps
|
||||
dist, _, _ = self.forward_train(
|
||||
cond,
|
||||
|
||||
@@ -21,6 +21,7 @@ class Gaussian_VisionMLP(nn.Module):
|
||||
transition_dim,
|
||||
horizon_steps,
|
||||
cond_dim,
|
||||
img_cond_steps=1,
|
||||
mlp_dims=[256, 256, 256],
|
||||
activation_type="Mish",
|
||||
residual_style=False,
|
||||
@@ -44,6 +45,8 @@ class Gaussian_VisionMLP(nn.Module):
|
||||
if augment:
|
||||
self.aug = RandomShiftsAug(pad=4)
|
||||
self.augment = augment
|
||||
self.num_img = num_img
|
||||
self.img_cond_steps = img_cond_steps
|
||||
if spatial_emb > 0:
|
||||
assert spatial_emb > 1, "this is the dimension"
|
||||
if num_img > 1:
|
||||
@@ -109,24 +112,31 @@ class Gaussian_VisionMLP(nn.Module):
|
||||
self.fixed_std = fixed_std
|
||||
self.learn_fixed_std = learn_fixed_std
|
||||
|
||||
def forward(self, x):
|
||||
B = len(x["state"])
|
||||
device = x["state"].device
|
||||
def forward(self, cond):
|
||||
B = len(cond["rgb"])
|
||||
device = cond["rgb"].device
|
||||
_, T_rgb, C, H, W = cond["rgb"].shape
|
||||
|
||||
# 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")
|
||||
# flatten history
|
||||
state = cond["state"].view(B, -1)
|
||||
|
||||
# Take recent images --- sometimes we want to use fewer img_cond_steps than cond_steps (e.g., 1 image but 3 prio)
|
||||
rgb = cond["rgb"][:, -self.img_cond_steps :]
|
||||
|
||||
# concatenate images in cond by channels
|
||||
if self.num_img > 1:
|
||||
rgb = rgb.reshape(B, T_rgb, self.num_img, 3, H, W)
|
||||
rgb = einops.rearrange(rgb, "b t n c h w -> b n (t 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"]
|
||||
rgb = einops.rearrange(rgb, "b t c h w -> b (t c) h w")
|
||||
|
||||
# convert rgb to float32 for augmentation
|
||||
rgb = rgb.float()
|
||||
|
||||
# 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.num_img > 1: # TODO: properly handle multiple images
|
||||
rgb1 = rgb[:, 0]
|
||||
rgb2 = rgb[:, 1]
|
||||
if self.augment:
|
||||
rgb1 = self.aug(rgb1)
|
||||
rgb2 = self.aug(rgb2)
|
||||
@@ -223,11 +233,15 @@ class Gaussian_MLP(nn.Module):
|
||||
self.fixed_std = fixed_std
|
||||
self.learn_fixed_std = learn_fixed_std
|
||||
|
||||
def forward(self, x):
|
||||
B = len(x)
|
||||
def forward(self, cond):
|
||||
B = len(cond["state"])
|
||||
device = cond["state"].device
|
||||
|
||||
# flatten history
|
||||
state = cond["state"].view(B, -1)
|
||||
|
||||
# mlp
|
||||
out_mean = self.mlp_mean(x)
|
||||
out_mean = self.mlp_mean(state)
|
||||
out_mean = torch.tanh(out_mean).view(
|
||||
B, self.horizon_steps * self.transition_dim
|
||||
) # tanh squashing in [-1, 1]
|
||||
@@ -238,9 +252,9 @@ class Gaussian_MLP(nn.Module):
|
||||
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
|
||||
out_scale = torch.ones_like(out_mean).to(device) * self.fixed_std
|
||||
else:
|
||||
out_logvar = self.mlp_logvar(x).view(
|
||||
out_logvar = self.mlp_logvar(state).view(
|
||||
B, self.horizon_steps * self.transition_dim
|
||||
)
|
||||
out_logvar = torch.clamp(out_logvar, self.logvar_min, self.logvar_max)
|
||||
|
||||
+10
-6
@@ -77,11 +77,15 @@ class GMM_MLP(nn.Module):
|
||||
use_layernorm=use_layernorm,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
B = len(x)
|
||||
def forward(self, cond):
|
||||
B = len(cond["state"])
|
||||
device = cond["state"].device
|
||||
|
||||
# flatten history
|
||||
state = cond["state"].view(B, -1)
|
||||
|
||||
# mlp
|
||||
out_mean = self.mlp_mean(x)
|
||||
out_mean = self.mlp_mean(state)
|
||||
out_mean = torch.tanh(out_mean).view(
|
||||
B, self.num_modes, self.horizon_steps * self.transition_dim
|
||||
) # tanh squashing in [-1, 1]
|
||||
@@ -92,15 +96,15 @@ class GMM_MLP(nn.Module):
|
||||
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
|
||||
out_scale = torch.ones_like(out_mean).to(device) * self.fixed_std
|
||||
else:
|
||||
out_logvar = self.mlp_logvar(x).view(
|
||||
out_logvar = self.mlp_logvar(state).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 = self.mlp_weights(state)
|
||||
out_weights = out_weights.view(B, self.num_modes)
|
||||
|
||||
return out_mean, out_scale, out_weights
|
||||
|
||||
@@ -67,3 +67,21 @@ class RandomShiftsAug:
|
||||
return nn.functional.grid_sample(
|
||||
x, grid, padding_mode="zeros", align_corners=False
|
||||
)
|
||||
|
||||
|
||||
# test random shift
|
||||
if __name__ == "__main__":
|
||||
from PIL import Image
|
||||
import requests
|
||||
import numpy as np
|
||||
|
||||
image_url = "https://rail.eecs.berkeley.edu/datasets/bridge_release/raw/bridge_data_v2/datacol2_toykitchen7/drawer_pnp/01/2023-04-19_09-18-15/raw/traj_group0/traj0/images0/im_30.jpg"
|
||||
image = Image.open(requests.get(image_url, stream=True).raw)
|
||||
image = image.resize((96, 96))
|
||||
|
||||
image = torch.tensor(np.array(image)).permute(2, 0, 1).unsqueeze(0).float()
|
||||
aug = RandomShiftsAug(pad=4)
|
||||
image_aug = aug(image)
|
||||
image_aug = image_aug.squeeze().permute(1, 2, 0).numpy()
|
||||
image_aug = Image.fromarray(image_aug.astype(np.uint8))
|
||||
image_aug.show()
|
||||
|
||||
+20
-16
@@ -70,13 +70,15 @@ class Gaussian_Transformer(nn.Module):
|
||||
)
|
||||
|
||||
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)
|
||||
B = len(cond["state"])
|
||||
device = cond["state"].device
|
||||
|
||||
# flatten history
|
||||
state = cond["state"].view(B, -1)
|
||||
|
||||
# input to transformer
|
||||
state = state.unsqueeze(1) # (B,1,cond_dim)
|
||||
out, _ = self.transformer(state) # (B,horizon,output_dim)
|
||||
|
||||
# use the first half of the output as mean
|
||||
out_mean = torch.tanh(out[:, :, : self.transition_dim])
|
||||
@@ -88,7 +90,7 @@ class Gaussian_Transformer(nn.Module):
|
||||
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
|
||||
out_scale = torch.ones_like(out_mean).to(device) * self.fixed_std
|
||||
else:
|
||||
out_logvar = out[:, :, self.transition_dim :]
|
||||
out_logvar = out_logvar.reshape(B, self.horizon_steps * self.transition_dim)
|
||||
@@ -164,14 +166,16 @@ class GMM_Transformer(nn.Module):
|
||||
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)
|
||||
B = len(cond["state"])
|
||||
device = cond["state"].device
|
||||
|
||||
# flatten history
|
||||
state = cond["state"].view(B, -1)
|
||||
|
||||
# input to transformer
|
||||
state = state.unsqueeze(1) # (B,1,cond_dim)
|
||||
out, out_prehead = self.transformer(
|
||||
cond
|
||||
state
|
||||
) # (B,horizon,output_dim), (B,horizon,emb_dim)
|
||||
|
||||
# use the first half of the output as mean
|
||||
@@ -190,7 +194,7 @@ class GMM_Transformer(nn.Module):
|
||||
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
|
||||
out_scale = torch.ones_like(out_mean).to(device) * self.fixed_std
|
||||
else:
|
||||
out_logvar = out[
|
||||
:, :, self.num_modes * self.transition_dim : -self.num_modes
|
||||
|
||||
+27
-21
@@ -24,7 +24,12 @@ class VitEncoderConfig:
|
||||
|
||||
|
||||
class VitEncoder(nn.Module):
|
||||
def __init__(self, obs_shape: List[int], cfg: VitEncoderConfig):
|
||||
def __init__(
|
||||
self,
|
||||
obs_shape: List[int],
|
||||
cfg: VitEncoderConfig,
|
||||
num_channel=3,
|
||||
):
|
||||
super().__init__()
|
||||
self.obs_shape = obs_shape
|
||||
self.cfg = cfg
|
||||
@@ -34,6 +39,7 @@ class VitEncoder(nn.Module):
|
||||
embed_norm=cfg.embed_norm,
|
||||
num_head=cfg.num_heads,
|
||||
depth=cfg.depth,
|
||||
num_channel=num_channel,
|
||||
)
|
||||
|
||||
self.num_patch = self.vit.num_patches
|
||||
@@ -50,9 +56,9 @@ class VitEncoder(nn.Module):
|
||||
|
||||
|
||||
class PatchEmbed1(nn.Module):
|
||||
def __init__(self, embed_dim):
|
||||
def __init__(self, embed_dim, num_channel=3):
|
||||
super().__init__()
|
||||
self.conv = nn.Conv2d(3, embed_dim, kernel_size=8, stride=8)
|
||||
self.conv = nn.Conv2d(num_channel, embed_dim, kernel_size=8, stride=8)
|
||||
|
||||
self.num_patch = 144
|
||||
self.patch_dim = embed_dim
|
||||
@@ -64,10 +70,10 @@ class PatchEmbed1(nn.Module):
|
||||
|
||||
|
||||
class PatchEmbed2(nn.Module):
|
||||
def __init__(self, embed_dim, use_norm):
|
||||
def __init__(self, embed_dim, use_norm, num_channel=3):
|
||||
super().__init__()
|
||||
layers = [
|
||||
nn.Conv2d(3, embed_dim, kernel_size=8, stride=4),
|
||||
nn.Conv2d(num_channel, 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),
|
||||
@@ -132,13 +138,23 @@ class TransformerLayer(nn.Module):
|
||||
|
||||
|
||||
class MinVit(nn.Module):
|
||||
def __init__(self, embed_style, embed_dim, embed_norm, num_head, depth):
|
||||
def __init__(
|
||||
self,
|
||||
embed_style,
|
||||
embed_dim,
|
||||
embed_norm,
|
||||
num_head,
|
||||
depth,
|
||||
num_channel=3,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
if embed_style == "embed1":
|
||||
self.patch_embed = PatchEmbed1(embed_dim)
|
||||
self.patch_embed = PatchEmbed1(embed_dim, num_channel=num_channel)
|
||||
elif embed_style == "embed2":
|
||||
self.patch_embed = PatchEmbed2(embed_dim, use_norm=embed_norm)
|
||||
self.patch_embed = PatchEmbed2(
|
||||
embed_dim, use_norm=embed_norm, num_channel=num_channel
|
||||
)
|
||||
else:
|
||||
assert False
|
||||
|
||||
@@ -217,20 +233,10 @@ def test_transformer_layer():
|
||||
|
||||
|
||||
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)
|
||||
obs_shape = [6, 96, 96]
|
||||
enc = VitEncoder([6, 96, 96], VitEncoderConfig())
|
||||
|
||||
print(enc)
|
||||
x = torch.rand(1, *cfg.obs_shape) * 255
|
||||
x = torch.rand(1, *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