squash commits

This commit is contained in:
allenzren
2024-09-11 21:09:17 -04:00
parent 8ce0aa1485
commit 2ddf63b8f5
200 changed files with 1240 additions and 1186 deletions
+70 -27
View File
@@ -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)
+8 -12
View File
@@ -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
View File
@@ -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,
+33 -19
View File
@@ -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
View File
@@ -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
+18
View File
@@ -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
View File
@@ -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
View File
@@ -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())