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())
|
||||
|
||||
@@ -8,7 +8,7 @@ Annotated DDIM/DDPM: https://nn.labml.ai/diffusion/stable_diffusion/sampler/ddpm
|
||||
|
||||
"""
|
||||
|
||||
from typing import Optional, Union
|
||||
from typing import Union
|
||||
import logging
|
||||
import torch
|
||||
from torch import nn
|
||||
@@ -17,13 +17,12 @@ import torch.nn.functional as F
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
from model.diffusion.sampling import (
|
||||
make_timesteps,
|
||||
extract,
|
||||
cosine_beta_schedule,
|
||||
)
|
||||
|
||||
from collections import namedtuple
|
||||
Sample = namedtuple("Sample", "trajectories values chains")
|
||||
Sample = namedtuple("Sample", "trajectories chains")
|
||||
|
||||
|
||||
class DiffusionModel(nn.Module):
|
||||
@@ -34,9 +33,7 @@ class DiffusionModel(nn.Module):
|
||||
horizon_steps,
|
||||
obs_dim,
|
||||
action_dim,
|
||||
transition_dim,
|
||||
network_path=None,
|
||||
cond_steps=1,
|
||||
device="cuda:0",
|
||||
# DDPM parameters
|
||||
denoising_steps=100,
|
||||
@@ -53,11 +50,9 @@ class DiffusionModel(nn.Module):
|
||||
self.horizon_steps = horizon_steps
|
||||
self.obs_dim = obs_dim
|
||||
self.action_dim = action_dim
|
||||
self.transition_dim = transition_dim
|
||||
self.denoising_steps = int(denoising_steps)
|
||||
self.denoised_clip_value = denoised_clip_value
|
||||
self.predict_epsilon = predict_epsilon
|
||||
self.cond_steps = cond_steps
|
||||
self.use_ddim = use_ddim
|
||||
self.ddim_steps = ddim_steps
|
||||
|
||||
@@ -216,52 +211,11 @@ class DiffusionModel(nn.Module):
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
cond: Optional[torch.Tensor],
|
||||
cond,
|
||||
return_chain=True,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Forward sampling through denoising steps.
|
||||
|
||||
Args:
|
||||
cond: (batch_size, horizon, transition_dim)
|
||||
return_chain: whether to return the chain of samples or only the final denoised sample
|
||||
Return:
|
||||
Sample: namedtuple with fields:
|
||||
trajectories: (batch_size, horizon_steps, transition_dim)
|
||||
values: (batch_size, )
|
||||
chain: (batch_size, denoising_steps + 1, horizon_steps, transition_dim)
|
||||
"""
|
||||
device = self.betas.device
|
||||
if isinstance(cond, dict):
|
||||
B = cond[list(cond.keys())[0]].shape[0]
|
||||
else:
|
||||
B = cond.shape[0]
|
||||
cond = cond[:, : self.cond_steps].reshape(B, -1)
|
||||
shape = (B, self.horizon_steps, self.transition_dim)
|
||||
|
||||
# Loop
|
||||
x = torch.randn(shape, device=device)
|
||||
chain = [x] if return_chain else None
|
||||
if self.use_ddim:
|
||||
t_all = self.ddim_t
|
||||
else:
|
||||
t_all = list(reversed(range(self.denoising_steps)))
|
||||
for i, t in enumerate(t_all):
|
||||
t_b = make_timesteps(B, t, device)
|
||||
index_b = make_timesteps(B, i, device)
|
||||
mu, logvar = self.p_mean_var(x=x, t=t_b, cond=cond, index=index_b)
|
||||
std = torch.exp(0.5 * logvar)
|
||||
|
||||
# no noise when t == 0
|
||||
noise = torch.randn_like(x)
|
||||
noise[t == 0] = 0
|
||||
x = mu + std * noise
|
||||
if return_chain:
|
||||
chain.append(x)
|
||||
if return_chain:
|
||||
chain = torch.stack(chain, dim=1)
|
||||
values = torch.zeros(len(x), device=x.device) # not considering the value for now
|
||||
return Sample(x, values, chain)
|
||||
raise NotImplementedError
|
||||
|
||||
# ---------- Supervised training ----------#
|
||||
|
||||
@@ -275,23 +229,18 @@ class DiffusionModel(nn.Module):
|
||||
def p_losses(
|
||||
self,
|
||||
x_start,
|
||||
obs_cond: Union[dict, torch.Tensor],
|
||||
cond: Union[dict, torch.Tensor],
|
||||
t,
|
||||
):
|
||||
"""
|
||||
If predicting epsilon: E_{t, x0, ε} [||ε - ε_θ(√α̅ₜx0 + √(1-α̅ₜ)ε, t)||²
|
||||
|
||||
Args:
|
||||
x_start: (batch_size, horizon_steps, transition_dim)
|
||||
obs_cond: dict with keys as step and value as observation
|
||||
x_start: (batch_size, horizon_steps, action_dim)
|
||||
cond: dict with keys as step and value as observation
|
||||
t: batch of integers
|
||||
"""
|
||||
device = x_start.device
|
||||
B = x_start.shape[0]
|
||||
if isinstance(obs_cond[0], dict):
|
||||
cond = obs_cond[0] # keep the dictionary and the network will extract img and prio
|
||||
else:
|
||||
cond = obs_cond[0].reshape(B, -1)
|
||||
|
||||
# Forward process
|
||||
noise = torch.randn_like(x_start, device=device)
|
||||
|
||||
@@ -40,19 +40,19 @@ class DIPODiffusion(DiffusionModel):
|
||||
# Whether to clamp sampled action between [-1, 1]
|
||||
self.clamp_action = clamp_action
|
||||
|
||||
# ---------- RL training ----------#
|
||||
|
||||
def loss_critic(self, obs, next_obs, actions, rewards, dones, gamma):
|
||||
|
||||
# get current Q-function
|
||||
actions_flat = torch.flatten(actions, start_dim=-2)
|
||||
current_q1, current_q2 = self.critic(obs, actions_flat)
|
||||
current_q1, current_q2 = self.critic(obs, actions)
|
||||
|
||||
# get next Q-function
|
||||
next_actions = self.forward(
|
||||
cond=next_obs,
|
||||
deterministic=False,
|
||||
) # in DiffusionModel, forward() has no gradient, which is desired here.
|
||||
next_actions_flat = torch.flatten(next_actions, start_dim=-2)
|
||||
next_q1, next_q2 = self.critic(next_obs, next_actions_flat)
|
||||
) # forward() has no gradient, which is desired here.
|
||||
next_q1, next_q2 = self.critic(next_obs, next_actions)
|
||||
next_q = torch.min(next_q1, next_q2)
|
||||
|
||||
# terminal state mask
|
||||
@@ -73,6 +73,8 @@ class DIPODiffusion(DiffusionModel):
|
||||
|
||||
return loss_critic
|
||||
|
||||
# ---------- Sampling ----------#``
|
||||
|
||||
# override
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
@@ -81,15 +83,10 @@ class DIPODiffusion(DiffusionModel):
|
||||
deterministic=False,
|
||||
):
|
||||
device = self.betas.device
|
||||
B = cond.shape[0]
|
||||
if isinstance(cond, dict):
|
||||
raise NotImplementedError("Not implemented for images")
|
||||
else:
|
||||
B = cond.shape[0]
|
||||
cond = cond[:, : self.cond_steps]
|
||||
B = len(cond["state"])
|
||||
|
||||
# Loop
|
||||
x = torch.randn((B, self.horizon_steps, self.transition_dim), device=device)
|
||||
x = torch.randn((B, self.horizon_steps, self.action_dim), device=device)
|
||||
t_all = list(reversed(range(self.denoising_steps)))
|
||||
for i, t in enumerate(t_all):
|
||||
t_b = make_timesteps(B, t, device)
|
||||
|
||||
@@ -41,19 +41,19 @@ class DQLDiffusion(DiffusionModel):
|
||||
# Whether to clamp sampled action between [-1, 1]
|
||||
self.clamp_action = clamp_action
|
||||
|
||||
# ---------- RL training ----------#
|
||||
|
||||
def loss_critic(self, obs, next_obs, actions, rewards, dones, gamma):
|
||||
|
||||
# get current Q-function
|
||||
actions_flat = torch.flatten(actions, start_dim=-2)
|
||||
current_q1, current_q2 = self.critic(obs, actions_flat)
|
||||
current_q1, current_q2 = self.critic(obs, actions)
|
||||
|
||||
# get next Q-function
|
||||
next_actions = self.forward(
|
||||
cond=next_obs,
|
||||
deterministic=False,
|
||||
) # in DiffusionModel, forward() has no gradient, which is desired here.
|
||||
next_actions_flat = torch.flatten(next_actions, start_dim=-2)
|
||||
next_q1, next_q2 = self.critic(next_obs, next_actions_flat)
|
||||
) # forward() has no gradient, which is desired here.
|
||||
next_q1, next_q2 = self.critic(next_obs, next_actions)
|
||||
next_q = torch.min(next_q1, next_q2)
|
||||
|
||||
# terminal state mask
|
||||
@@ -75,7 +75,7 @@ class DQLDiffusion(DiffusionModel):
|
||||
return loss_critic
|
||||
|
||||
def loss_actor(self, obs, actions, q1, q2, eta):
|
||||
bc_loss = self.loss(actions, {0: obs})
|
||||
bc_loss = self.loss(actions, obs)
|
||||
if np.random.uniform() > 0.5:
|
||||
q_loss = -q1.mean() / q2.abs().mean().detach()
|
||||
else:
|
||||
@@ -83,6 +83,8 @@ class DQLDiffusion(DiffusionModel):
|
||||
actor_loss = bc_loss + eta * q_loss
|
||||
return actor_loss
|
||||
|
||||
# ---------- Sampling ----------#``
|
||||
|
||||
# override
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
@@ -91,15 +93,10 @@ class DQLDiffusion(DiffusionModel):
|
||||
deterministic=False,
|
||||
):
|
||||
device = self.betas.device
|
||||
B = cond.shape[0]
|
||||
if isinstance(cond, dict):
|
||||
raise NotImplementedError("Not implemented for images")
|
||||
else:
|
||||
B = cond.shape[0]
|
||||
cond = cond[:, : self.cond_steps]
|
||||
B = len(cond["state"])
|
||||
|
||||
# Loop
|
||||
x = torch.randn((B, self.horizon_steps, self.transition_dim), device=device)
|
||||
x = torch.randn((B, self.horizon_steps, self.action_dim), device=device)
|
||||
t_all = list(reversed(range(self.denoising_steps)))
|
||||
for i, t in enumerate(t_all):
|
||||
t_b = make_timesteps(B, t, device)
|
||||
@@ -136,15 +133,10 @@ class DQLDiffusion(DiffusionModel):
|
||||
Differentiable forward pass used in actor training.
|
||||
"""
|
||||
device = self.betas.device
|
||||
B = cond.shape[0]
|
||||
if isinstance(cond, dict):
|
||||
raise NotImplementedError("Not implemented for images")
|
||||
else:
|
||||
B = cond.shape[0]
|
||||
cond = cond[:, : self.cond_steps]
|
||||
B = len(cond["state"])
|
||||
|
||||
# Loop
|
||||
x = torch.randn((B, self.horizon_steps, self.transition_dim), device=device)
|
||||
x = torch.randn((B, self.horizon_steps, self.action_dim), device=device)
|
||||
t_all = list(reversed(range(self.denoising_steps)))
|
||||
for i, t in enumerate(t_all):
|
||||
t_b = make_timesteps(B, t, device)
|
||||
|
||||
@@ -42,12 +42,13 @@ class IDQLDiffusion(RWRDiffusion):
|
||||
# assign actor
|
||||
self.actor = self.network
|
||||
|
||||
# ---------- RL training ----------#
|
||||
|
||||
def compute_advantages(self, obs, actions):
|
||||
|
||||
# get current Q-function
|
||||
actions_flat = torch.flatten(actions, start_dim=-2)
|
||||
with torch.no_grad(): # no gradients for q-function when we update value function
|
||||
current_q1, current_q2 = self.target_q(obs, actions_flat)
|
||||
# get current Q-function, stop gradient
|
||||
with torch.no_grad():
|
||||
current_q1, current_q2 = self.target_q(obs, actions)
|
||||
q = torch.min(current_q1, current_q2)
|
||||
|
||||
# get the current V-function
|
||||
@@ -59,7 +60,6 @@ class IDQLDiffusion(RWRDiffusion):
|
||||
return adv
|
||||
|
||||
def loss_critic_v(self, obs, actions):
|
||||
|
||||
adv = self.compute_advantages(obs, actions)
|
||||
|
||||
# get the value loss
|
||||
@@ -70,11 +70,10 @@ class IDQLDiffusion(RWRDiffusion):
|
||||
def loss_critic_q(self, obs, next_obs, actions, rewards, dones, gamma):
|
||||
|
||||
# get current Q-function
|
||||
actions_flat = torch.flatten(actions, start_dim=-2)
|
||||
current_q1, current_q2 = self.critic_q(obs, actions_flat)
|
||||
current_q1, current_q2 = self.critic_q(obs, actions)
|
||||
|
||||
# get the next V-function
|
||||
with torch.no_grad(): # no gradients for value function when we update q function
|
||||
# get the next V-function, stop gradient
|
||||
with torch.no_grad():
|
||||
next_v = self.critic_v(next_obs)
|
||||
|
||||
# terminal state mask
|
||||
@@ -98,8 +97,35 @@ class IDQLDiffusion(RWRDiffusion):
|
||||
def update_target_critic(self, tau):
|
||||
soft_update(self.target_q, self.critic_q, tau)
|
||||
|
||||
# override
|
||||
def p_losses(
|
||||
self,
|
||||
x_start,
|
||||
cond,
|
||||
t,
|
||||
):
|
||||
device = x_start.device
|
||||
|
||||
# Forward process
|
||||
noise = torch.randn_like(x_start, device=device)
|
||||
x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise)
|
||||
|
||||
# Predict
|
||||
x_recon = self.network(x_noisy, t, cond=cond)
|
||||
|
||||
# Loss with mask
|
||||
if self.predict_epsilon:
|
||||
loss = F.mse_loss(x_recon, noise, reduction="none")
|
||||
else:
|
||||
loss = F.mse_loss(x_recon, x_start, reduction="none")
|
||||
loss = einops.reduce(loss, "b h d -> b", "mean")
|
||||
return loss.mean()
|
||||
|
||||
# ---------- Sampling ----------#``
|
||||
|
||||
# override
|
||||
@torch.no_grad()
|
||||
def forward( # override
|
||||
def forward(
|
||||
self,
|
||||
cond,
|
||||
deterministic=False,
|
||||
@@ -107,23 +133,23 @@ class IDQLDiffusion(RWRDiffusion):
|
||||
critic_hyperparam=0.7, # sampling weight for implicit policy
|
||||
use_expectile_exploration=True,
|
||||
):
|
||||
"""assume state-only, no rgb in cond"""
|
||||
# repeat obs num_sample times along dim 0
|
||||
cond_shape_repeat_dims = tuple(1 for _ in cond.shape)
|
||||
B, T, D = cond.shape
|
||||
cond_shape_repeat_dims = tuple(1 for _ in cond["state"].shape)
|
||||
B, T, D = cond["state"].shape
|
||||
S = num_sample
|
||||
cond_repeat = cond[None].repeat(num_sample, *cond_shape_repeat_dims)
|
||||
cond_repeat = cond["state"][None].repeat(num_sample, *cond_shape_repeat_dims)
|
||||
cond_repeat = cond_repeat.view(-1, T, D) # [B*S, T, D]
|
||||
|
||||
# for eval, use less noisy samples --- there is still DDPM noise, but final action uses small min_sampling_std
|
||||
samples = super(IDQLDiffusion, self).forward(
|
||||
cond_repeat,
|
||||
{"state": cond_repeat},
|
||||
deterministic=deterministic,
|
||||
)
|
||||
_, H, A = samples.shape
|
||||
|
||||
# get current Q-function
|
||||
actions_flat = torch.flatten(samples, start_dim=-2)
|
||||
current_q1, current_q2 = self.target_q(cond_repeat, actions_flat)
|
||||
current_q1, current_q2 = self.target_q({"state": cond_repeat}, samples)
|
||||
q = torch.min(current_q1, current_q2)
|
||||
q = q.view(S, B)
|
||||
|
||||
@@ -141,7 +167,7 @@ class IDQLDiffusion(RWRDiffusion):
|
||||
# Sample as an implicit policy for exploration
|
||||
else:
|
||||
# get the current value function for probabilistic exploration
|
||||
current_v = self.critic_v(cond_repeat)
|
||||
current_v = self.critic_v({"state": cond_repeat})
|
||||
v = current_v.view(S, B)
|
||||
adv = q - v
|
||||
|
||||
@@ -164,34 +190,3 @@ class IDQLDiffusion(RWRDiffusion):
|
||||
# squeeze dummy dimension
|
||||
samples = samples_best[0]
|
||||
return samples
|
||||
|
||||
# override
|
||||
def p_losses(
|
||||
self,
|
||||
x_start,
|
||||
obs_cond,
|
||||
t,
|
||||
):
|
||||
device = x_start.device
|
||||
B, T, D = x_start.shape
|
||||
|
||||
# handle different ways of passing observation
|
||||
if isinstance(obs_cond[0], dict):
|
||||
cond = obs_cond[0]
|
||||
else:
|
||||
cond = obs_cond.reshape(B, -1)
|
||||
|
||||
# Forward process
|
||||
noise = torch.randn_like(x_start, device=device)
|
||||
x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise)
|
||||
|
||||
# Predict
|
||||
x_recon = self.network(x_noisy, t, cond=cond)
|
||||
|
||||
# Loss with mask
|
||||
if self.predict_epsilon:
|
||||
loss = F.mse_loss(x_recon, noise, reduction="none")
|
||||
else:
|
||||
loss = F.mse_loss(x_recon, x_start, reduction="none")
|
||||
loss = einops.reduce(loss, "b h d -> b", "mean")
|
||||
return loss.mean()
|
||||
|
||||
@@ -1,6 +1,15 @@
|
||||
"""
|
||||
DPPO: Diffusion Policy Policy Optimization.
|
||||
|
||||
K: number of denoising steps
|
||||
To: observation sequence length
|
||||
Ta: action chunk size
|
||||
Do: observation dimension
|
||||
Da: action dimension
|
||||
|
||||
C: image channels
|
||||
H, W: image height and width
|
||||
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
@@ -60,13 +69,15 @@ class PPODiffusion(VPGDiffusion):
|
||||
"""
|
||||
PPO loss
|
||||
|
||||
obs: (B, obs_step, obs_dim)
|
||||
chains: (B, num_denoising_step+1, horizon_step, action_dim)
|
||||
obs: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
rgb: (B, To, C, H, W)
|
||||
chains: (B, K+1, Ta, Da)
|
||||
returns: (B, )
|
||||
values: (B, )
|
||||
advantages: (B,)
|
||||
oldlogprobs: (B, num_denoising_step, horizon_step, action_dim)
|
||||
use_bc_loss: add BC regularization loss
|
||||
oldlogprobs: (B, K, Ta, Da)
|
||||
use_bc_loss: whether to add BC regularization loss
|
||||
reward_horizon: action horizon that backpropagates gradient
|
||||
"""
|
||||
# Get new logprobs for denoising steps from T-1 to 0 - entropy is fixed fod diffusion
|
||||
|
||||
@@ -3,6 +3,11 @@ Diffusion policy gradient with exact likelihood estimation.
|
||||
|
||||
Based on score_sde_pytorch https://github.com/yang-song/score_sde_pytorch
|
||||
|
||||
To: observation sequence length
|
||||
Ta: action chunk size
|
||||
Do: observation dimension
|
||||
Da: action dimension
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
@@ -52,13 +57,12 @@ class PPOExactDiffusion(PPODiffusion):
|
||||
num_epsilon=sde_num_epsilon,
|
||||
)
|
||||
|
||||
def get_exact_logprobs(self, obs, samples):
|
||||
def get_exact_logprobs(self, cond, samples):
|
||||
"""Use torchdiffeq
|
||||
|
||||
samples: B x horizon x transition_dim
|
||||
samples: (B x Ta x Da)
|
||||
"""
|
||||
# TODO: image input
|
||||
cond = obs.reshape(-1, self.obs_dim)
|
||||
return self.likelihood_fn(
|
||||
self.actor,
|
||||
self.actor_ft,
|
||||
@@ -79,6 +83,17 @@ class PPOExactDiffusion(PPODiffusion):
|
||||
use_bc_loss=False,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
PPO loss
|
||||
|
||||
obs: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
samples: (B, Ta, Da)
|
||||
returns: (B, )
|
||||
values: (B, )
|
||||
advantages: (B,)
|
||||
oldlogprobs: (B, )
|
||||
"""
|
||||
# Get new logprobs for final x
|
||||
newlogprobs = self.get_exact_logprobs(obs, samples)
|
||||
newlogprobs = newlogprobs.clamp(min=-5, max=2)
|
||||
|
||||
@@ -39,11 +39,12 @@ class QSMDiffusion(RWRDiffusion):
|
||||
# assign actor
|
||||
self.actor = self.network
|
||||
|
||||
# ---------- RL training ----------#
|
||||
|
||||
def loss_actor(self, obs, actions, q_grad_coeff):
|
||||
x_start = actions
|
||||
device = x_start.device
|
||||
B, T, D = x_start.shape
|
||||
cond = obs.reshape(B, -1)
|
||||
B = len(x_start)
|
||||
|
||||
# Forward process
|
||||
noise = torch.randn_like(x_start, device=device)
|
||||
@@ -53,39 +54,35 @@ class QSMDiffusion(RWRDiffusion):
|
||||
x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise)
|
||||
|
||||
# get current value for noisy actions as the code does --- the algorthm block in the paper is wrong, it says using a_t, the final denoised action
|
||||
x_noisy_flat = torch.flatten(x_noisy, start_dim=-2)
|
||||
x_noisy_flat.requires_grad_(True)
|
||||
current_q1, current_q2 = self.critic_q(obs, x_noisy_flat)
|
||||
# x_noisy_flat = torch.flatten(x_noisy, start_dim=-2)
|
||||
x_noisy.requires_grad_(True)
|
||||
current_q1, current_q2 = self.critic_q(obs, x_noisy)
|
||||
|
||||
# Compute dQ/da|a=noise_actions
|
||||
gradient_q1 = torch.autograd.grad(current_q1.sum(), x_noisy_flat)[0]
|
||||
gradient_q2 = torch.autograd.grad(current_q2.sum(), x_noisy_flat)[0]
|
||||
gradient_q1 = torch.autograd.grad(current_q1.sum(), x_noisy)[0]
|
||||
gradient_q2 = torch.autograd.grad(current_q2.sum(), x_noisy)[0]
|
||||
gradient_q = torch.stack((gradient_q1, gradient_q2), 0).mean(0).detach()
|
||||
|
||||
# Predict noise from noisy actions
|
||||
x_recon = self.network(x_noisy, t, cond=cond)
|
||||
x_recon = torch.flatten(x_recon, start_dim=-2)
|
||||
x_recon = self.network(x_noisy, t, cond=obs)
|
||||
|
||||
# Loss with mask - align predicted noise with critic gradient of noisy actions
|
||||
# Note: the gradient of mu wrt. epsilon has a negative sign
|
||||
loss = F.mse_loss(-x_recon, q_grad_coeff * gradient_q, reduction="none").mean()
|
||||
|
||||
return loss
|
||||
|
||||
def loss_critic(self, obs, next_obs, actions, rewards, dones, gamma):
|
||||
|
||||
# get current Q-function
|
||||
actions_flat = torch.flatten(actions, start_dim=-2)
|
||||
current_q1, current_q2 = self.critic_q(obs, actions_flat)
|
||||
current_q1, current_q2 = self.critic_q(obs, actions)
|
||||
|
||||
# get next Q-function - with noise, same as QSM https://github.com/Alescontrela/score_matching_rl/blob/f02a21969b17e322eb229ceb2b0f5a9111b1b968/jaxrl5/agents/score_matching/score_matching_learner.py#L193
|
||||
next_actions = self.forward(
|
||||
cond=next_obs,
|
||||
deterministic=False,
|
||||
) # in DiffusionModel, forward() has no gradient, which is desired here.
|
||||
next_actions_flat = torch.flatten(next_actions, start_dim=-2)
|
||||
) # forward() has no gradient, which is desired here.
|
||||
with torch.no_grad():
|
||||
next_q1, next_q2 = self.target_q(next_obs, next_actions_flat)
|
||||
next_q1, next_q2 = self.target_q(next_obs, next_actions)
|
||||
next_q = torch.min(next_q1, next_q2)
|
||||
|
||||
# terminal state mask
|
||||
|
||||
@@ -43,18 +43,11 @@ class RWRDiffusion(DiffusionModel):
|
||||
def p_losses(
|
||||
self,
|
||||
x_start,
|
||||
obs_cond,
|
||||
cond,
|
||||
rewards,
|
||||
t,
|
||||
):
|
||||
device = x_start.device
|
||||
B, T, D = x_start.shape
|
||||
|
||||
# handle different ways of passing observation
|
||||
if isinstance(obs_cond[0], dict):
|
||||
cond = obs_cond[0]
|
||||
else:
|
||||
cond = obs_cond.reshape(B, -1)
|
||||
|
||||
# Forward process
|
||||
noise = torch.randn_like(x_start, device=device)
|
||||
@@ -79,7 +72,7 @@ class RWRDiffusion(DiffusionModel):
|
||||
self,
|
||||
x,
|
||||
t,
|
||||
cond=None,
|
||||
cond,
|
||||
):
|
||||
noise = self.network(x, t, cond=cond)
|
||||
|
||||
@@ -116,15 +109,10 @@ class RWRDiffusion(DiffusionModel):
|
||||
deterministic=False,
|
||||
):
|
||||
device = self.betas.device
|
||||
B = cond.shape[0]
|
||||
if isinstance(cond, dict):
|
||||
raise NotImplementedError("Not implemented for images")
|
||||
else:
|
||||
B = cond.shape[0]
|
||||
cond = cond[:, : self.cond_steps]
|
||||
B = len(cond["state"])
|
||||
|
||||
# Loop
|
||||
x = torch.randn((B, self.horizon_steps, self.transition_dim), device=device)
|
||||
x = torch.randn((B, self.horizon_steps, self.action_dim), device=device)
|
||||
t_all = list(reversed(range(self.denoising_steps)))
|
||||
for i, t in enumerate(t_all):
|
||||
t_b = make_timesteps(B, t, device)
|
||||
|
||||
@@ -1,14 +1,20 @@
|
||||
"""
|
||||
Policy gradient with diffusion policy.
|
||||
Policy gradient with diffusion policy. VPG: vanilla policy gradient
|
||||
|
||||
VPG: vanilla policy gradient
|
||||
K: number of denoising steps
|
||||
To: observation sequence length
|
||||
Ta: action chunk size
|
||||
Do: observation dimension
|
||||
Da: action dimension
|
||||
|
||||
C: image channels
|
||||
H, W: image height and width
|
||||
|
||||
"""
|
||||
|
||||
import copy
|
||||
import torch
|
||||
import logging
|
||||
import einops
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
import torch.nn.functional as F
|
||||
@@ -106,7 +112,11 @@ class VPGDiffusion(DiffusionModel):
|
||||
# ---------- Sampling ----------#
|
||||
|
||||
def step(self):
|
||||
"""Update min_sampling_denoising_std annealing and fine-tuning denoising steps annealing. Both not used currently"""
|
||||
"""
|
||||
Anneal min_sampling_denoising_std and fine-tuning denoising steps
|
||||
|
||||
Current configs do not apply annealing
|
||||
"""
|
||||
# anneal min_sampling_denoising_std
|
||||
if type(self.min_sampling_denoising_std) is not float:
|
||||
self.min_sampling_denoising_std.step()
|
||||
@@ -142,7 +152,7 @@ class VPGDiffusion(DiffusionModel):
|
||||
self,
|
||||
x,
|
||||
t,
|
||||
cond=None,
|
||||
cond,
|
||||
index=None,
|
||||
use_base_policy=False,
|
||||
deterministic=False,
|
||||
@@ -160,12 +170,8 @@ class VPGDiffusion(DiffusionModel):
|
||||
|
||||
# overwrite noise for fine-tuning steps
|
||||
if len(ft_indices) > 0:
|
||||
if cond is not None:
|
||||
if isinstance(cond, dict):
|
||||
cond = {key: cond[key][ft_indices] for key in cond}
|
||||
else:
|
||||
cond = cond[ft_indices]
|
||||
noise_ft = actor(x[ft_indices], t[ft_indices], cond=cond)
|
||||
cond_ft = {key: cond[key][ft_indices] for key in cond}
|
||||
noise_ft = actor(x[ft_indices], t[ft_indices], cond=cond_ft)
|
||||
noise[ft_indices] = noise_ft
|
||||
|
||||
# Predict x_0
|
||||
@@ -208,7 +214,8 @@ class VPGDiffusion(DiffusionModel):
|
||||
if deterministic:
|
||||
etas = torch.zeros((x.shape[0], 1, 1)).to(x.device)
|
||||
else:
|
||||
etas = self.eta(cond).unsqueeze(1) # B x 1 x (transition_dim or 1)
|
||||
# TODO: eta cond
|
||||
etas = self.eta(cond).unsqueeze(1) # B x 1 x (Da or 1)
|
||||
sigma = (
|
||||
etas
|
||||
* ((1 - alpha_prev) / (1 - alpha) * (1 - alpha / alpha_prev)) ** 0.5
|
||||
@@ -242,30 +249,26 @@ class VPGDiffusion(DiffusionModel):
|
||||
Forward pass for sampling actions.
|
||||
|
||||
Args:
|
||||
cond: (batch_size, obs_step, obs_dim)
|
||||
deterministic: whether to sample deterministically
|
||||
return_chain: whether to return the chain of samples
|
||||
use_base_policy: whether to use the base policy instead
|
||||
cond: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
rgb: (B, To, C, H, W)
|
||||
deterministic: If true, then std=0 with DDIM, or with DDPM, use normal schedule (instead of clipping at a higher value)
|
||||
return_chain: whether to return the entire chain of denoised actions
|
||||
use_base_policy: whether to use the frozen pre-trained policy instead
|
||||
Return:
|
||||
Sample: namedtuple with fields:
|
||||
trajectories: (batch_size, horizon_steps, transition_dim)
|
||||
values: (batch_size, )
|
||||
chain: (batch_size, denoising_steps + 1, horizon_steps, transition_dim)
|
||||
trajectories: (B, Ta, Da)
|
||||
chain: (B, K + 1, Ta, Da)
|
||||
"""
|
||||
device = self.betas.device
|
||||
if isinstance(cond, dict):
|
||||
B = cond["state"].shape[0]
|
||||
cond["state"] = cond["state"][:, : self.cond_steps]
|
||||
cond["rgb"] = cond["rgb"][:, : self.cond_steps]
|
||||
else:
|
||||
B = cond.shape[0]
|
||||
cond = cond[:, : self.cond_steps]
|
||||
sample_data = cond["state"] if "state" in cond else cond["rgb"]
|
||||
B = len(sample_data)
|
||||
|
||||
# Get updated minimum sampling denoising std
|
||||
min_sampling_denoising_std = self.get_min_sampling_denoising_std()
|
||||
|
||||
# Loop
|
||||
x = torch.randn((B, self.horizon_steps, self.transition_dim), device=device)
|
||||
x = torch.randn((B, self.horizon_steps, self.action_dim), device=device)
|
||||
if self.use_ddim:
|
||||
t_all = self.ddim_t
|
||||
else:
|
||||
@@ -317,17 +320,16 @@ class VPGDiffusion(DiffusionModel):
|
||||
self.ddim_steps - self.ft_denoising_steps - 1
|
||||
):
|
||||
chain.append(x)
|
||||
values = torch.zeros(len(x), device=x.device)
|
||||
|
||||
if return_chain:
|
||||
chain = torch.stack(chain, dim=1)
|
||||
return Sample(x, values, chain)
|
||||
return Sample(x, chain)
|
||||
|
||||
# ---------- RL training ----------#
|
||||
|
||||
def get_logprobs(
|
||||
self,
|
||||
obs,
|
||||
cond,
|
||||
chains,
|
||||
get_ent: bool = False,
|
||||
use_base_policy: bool = False,
|
||||
@@ -336,35 +338,25 @@ class VPGDiffusion(DiffusionModel):
|
||||
Calculating the logprobs of the entire chain of denoised actions.
|
||||
|
||||
Args:
|
||||
obs: (B, obs_step, obs_dim)
|
||||
chains: (B, num_denoising_step+1, horizon_step, action_dim)
|
||||
get_ent: flag for returning entropy
|
||||
use_base_policy: flag for using base policy
|
||||
cond: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
rgb: (B, To, C, H, W)
|
||||
chains: (B, K+1, Ta, Da)
|
||||
get_ent: flag for returning entropy
|
||||
use_base_policy: flag for using base policy
|
||||
|
||||
Returns:
|
||||
logprobs: (B x num_denoising_steps, horizon_step, action_dim)
|
||||
entropy (if get_ent=True): (B x num_denoising_steps, horizon_step)
|
||||
logprobs: (B x K, Ta, Da)
|
||||
entropy (if get_ent=True): (B x K, Ta)
|
||||
"""
|
||||
# Repeat obs conditioning for denoising_steps
|
||||
if isinstance(obs, dict):
|
||||
obs = {
|
||||
key: obs[key]
|
||||
.unsqueeze(1)
|
||||
.repeat(1, self.ft_denoising_steps, *(1,) * (obs[key].ndim - 1))
|
||||
for key in obs
|
||||
}
|
||||
else:
|
||||
obs = einops.repeat(obs, "b h d -> b t h d", t=self.ft_denoising_steps)
|
||||
|
||||
# flatten the first two dimensions
|
||||
if isinstance(obs, dict):
|
||||
cond = obs
|
||||
for key in cond:
|
||||
cond[key] = einops.rearrange(cond[key], "b t ... -> (b t) ...")
|
||||
cond[key] = cond[key][:, : self.cond_steps]
|
||||
else:
|
||||
cond = einops.rearrange(obs, "b t h d -> (b t) h d")
|
||||
cond = cond[:, : self.cond_steps]
|
||||
# Repeat cond for denoising_steps, flatten batch and time dimensions
|
||||
cond = {
|
||||
key: cond[key]
|
||||
.unsqueeze(1)
|
||||
.repeat(1, self.ft_denoising_steps, *(1,) * (cond[key].ndim - 1))
|
||||
.flatten(start_dim=0, end_dim=1)
|
||||
for key in cond
|
||||
} # less memory usage than einops?
|
||||
|
||||
# Repeat t for batch dim, keep it 1-dim
|
||||
if self.use_ddim:
|
||||
@@ -393,8 +385,8 @@ class VPGDiffusion(DiffusionModel):
|
||||
chains_next = chains[:, 1:]
|
||||
|
||||
# Flatten first two dimensions
|
||||
chains_prev = chains_prev.reshape(-1, self.horizon_steps, self.transition_dim)
|
||||
chains_next = chains_next.reshape(-1, self.horizon_steps, self.transition_dim)
|
||||
chains_prev = chains_prev.reshape(-1, self.horizon_steps, self.action_dim)
|
||||
chains_next = chains_next.reshape(-1, self.horizon_steps, self.action_dim)
|
||||
|
||||
# Forward pass with previous chains
|
||||
next_mean, logvar, eta = self.p_mean_var(
|
||||
@@ -414,44 +406,35 @@ class VPGDiffusion(DiffusionModel):
|
||||
return log_prob, eta
|
||||
return log_prob
|
||||
|
||||
def loss(self, obs, chains, reward):
|
||||
def loss(self, cond, chains, reward):
|
||||
"""
|
||||
REINFORCE loss. Not used right now.
|
||||
|
||||
Args:
|
||||
obs: (n_steps, n_envs, obs_dim)
|
||||
chains: (n_steps, n_envs, num_denoising_step+1, horizon_step, action_dim)
|
||||
reward (to go): (n_steps, n_envs)
|
||||
cond: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
rgb: (B, To, C, H, W)
|
||||
chains: (B, K+1, Ta, Da)
|
||||
reward (to go): (b,)
|
||||
"""
|
||||
if torch.is_tensor(reward):
|
||||
assert not reward.requires_grad
|
||||
|
||||
n_steps, n_envs, _ = obs.shape
|
||||
|
||||
# Flatten first two dimensions
|
||||
obs = einops.rearrange(obs, "s e d -> (s e) d")
|
||||
chains = einops.rearrange(chains, "s e t h d -> (s e) t h d")
|
||||
reward = reward.reshape(-1)
|
||||
|
||||
# Get advantage
|
||||
with torch.no_grad():
|
||||
value = self.critic(obs).squeeze()
|
||||
value = self.critic(cond).squeeze()
|
||||
advantage = reward - value
|
||||
|
||||
# Get logprobs for denoising steps from T-1 to 0
|
||||
logprobs, eta = self.get_logprobs(obs, chains, get_ent=True)
|
||||
# (n_steps x n_envs x denoising_steps) x horizon_steps x (obs_dim+action_dim)
|
||||
logprobs, eta = self.get_logprobs(cond, chains, get_ent=True)
|
||||
# (n_steps x n_envs x K) x Ta x (Do+Da)
|
||||
|
||||
# Ignore obs dimension, and then sum over action dimension
|
||||
logprobs = logprobs[:, :, : self.action_dim].sum(-1)
|
||||
# -> (n_steps x n_envs x denoising_steps) x horizon_steps
|
||||
# -> (n_steps x n_envs x K) x Ta
|
||||
|
||||
# -> (n_steps x n_envs) x K x Ta
|
||||
logprobs = logprobs.reshape((-1, self.denoising_steps, self.horizon_steps))
|
||||
|
||||
# -> (n_steps x n_envs) x denoising_steps x horizon_steps
|
||||
logprobs = logprobs.reshape(
|
||||
(n_steps * n_envs, self.denoising_steps, self.horizon_steps)
|
||||
)
|
||||
# Sum/avg over denoising steps
|
||||
logprobs = logprobs.mean(-2) # -> (n_steps x n_envs) x horizon_steps
|
||||
logprobs = logprobs.mean(-2) # -> (n_steps x n_envs) x Ta
|
||||
|
||||
# Sum/avg over horizon steps
|
||||
logprobs = logprobs.mean(-1) # -> (n_steps x n_envs)
|
||||
@@ -460,6 +443,6 @@ class VPGDiffusion(DiffusionModel):
|
||||
loss_actor = torch.mean(-logprobs * advantage)
|
||||
|
||||
# Train critic to predict state value
|
||||
pred = self.critic(obs).squeeze()
|
||||
pred = self.critic(cond).squeeze()
|
||||
loss_critic = F.mse_loss(pred, reward)
|
||||
return loss_actor, loss_critic, eta
|
||||
|
||||
+24
-22
@@ -28,14 +28,11 @@ class EtaFixed(torch.nn.Module):
|
||||
torch.tensor([2 * (base_eta - min_eta) / (max_eta - min_eta) - 1])
|
||||
)
|
||||
|
||||
def __call__(self, x):
|
||||
def __call__(self, cond):
|
||||
"""Match input batch size, but do not depend on input"""
|
||||
if isinstance(x, dict):
|
||||
B = x["state"].shape[0]
|
||||
device = x["state"].device
|
||||
else:
|
||||
B = x.size(0)
|
||||
device = x.device
|
||||
sample_data = cond["state"] if "state" in cond else cond["rgb"]
|
||||
B = len(sample_data)
|
||||
device = sample_data.device
|
||||
eta_normalized = torch.tanh(self.eta_logit)
|
||||
|
||||
# map to min and max from [-1, 1]
|
||||
@@ -64,14 +61,11 @@ class EtaAction(torch.nn.Module):
|
||||
self.min = min_eta
|
||||
self.max = max_eta
|
||||
|
||||
def __call__(self, x):
|
||||
def __call__(self, cond):
|
||||
"""Match input batch size, but do not depend on input"""
|
||||
if isinstance(x, dict):
|
||||
B = x["state"].shape[0]
|
||||
device = x["state"].device
|
||||
else:
|
||||
B = x.size(0)
|
||||
device = x.device
|
||||
sample_data = cond["state"] if "state" in cond else cond["rgb"]
|
||||
B = len(sample_data)
|
||||
device = sample_data.device
|
||||
eta_normalized = torch.tanh(self.eta_logit)
|
||||
|
||||
# map to min and max from [-1, 1]
|
||||
@@ -109,13 +103,17 @@ class EtaState(torch.nn.Module):
|
||||
torch.nn.init.xavier_normal_(m.weight, gain=gain)
|
||||
m.bias.data.fill_(0)
|
||||
|
||||
def __call__(self, x):
|
||||
if isinstance(x, dict):
|
||||
def __call__(self, cond):
|
||||
if "rgb" in cond:
|
||||
raise NotImplementedError(
|
||||
"State-based eta not implemented for image-based training!"
|
||||
)
|
||||
x = x.view(x.size(0), -1)
|
||||
eta_res = self.mlp_res(x)
|
||||
# flatten history
|
||||
B = len(cond["state"])
|
||||
state = cond["state"].view(B, -1)
|
||||
|
||||
# forward pass
|
||||
eta_res = self.mlp_res(state)
|
||||
eta_res = torch.tanh(eta_res) # [-1, 1]
|
||||
eta = eta_res + self.base # [0, 2]
|
||||
return torch.clamp(eta, self.min_res + self.base, self.max_res + self.base)
|
||||
@@ -152,13 +150,17 @@ class EtaStateAction(torch.nn.Module):
|
||||
torch.nn.init.xavier_normal_(m.weight, gain=gain)
|
||||
m.bias.data.fill_(0)
|
||||
|
||||
def __call__(self, x):
|
||||
if isinstance(x, dict):
|
||||
def __call__(self, cond):
|
||||
if "rgb" in cond:
|
||||
raise NotImplementedError(
|
||||
"State-action-based eta not implemented for image-based training!"
|
||||
)
|
||||
x = x.view(x.size(0), -1)
|
||||
eta_res = self.mlp_res(x)
|
||||
# flatten history
|
||||
B = len(cond["state"])
|
||||
state = cond["state"].view(B, -1)
|
||||
|
||||
# forward pass
|
||||
eta_res = self.mlp_res(state)
|
||||
eta_res = torch.tanh(eta_res) # [-1, 1]
|
||||
eta = eta_res + self.base
|
||||
return torch.clamp(eta, self.min_res + self.base, self.max_res + self.base)
|
||||
|
||||
@@ -93,11 +93,12 @@ def get_likelihood_fn(
|
||||
"""Compute an unbiased estimate to the log-likelihood in bits/dim.
|
||||
|
||||
Args:
|
||||
model: A score model.
|
||||
data: A PyTorch tensor. B x horizon x transition_dim
|
||||
cond: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
data: (B x Ta x Da)
|
||||
|
||||
Returns:
|
||||
logprob: B
|
||||
logprob: (B,)
|
||||
"""
|
||||
shape = data.shape
|
||||
B, H, A = shape
|
||||
@@ -118,7 +119,9 @@ def get_likelihood_fn(
|
||||
raise NotImplementedError(f"Hutchinson type {hutchinson_type} unknown.")
|
||||
|
||||
# repeat for expectation
|
||||
cond_eps = cond.repeat_interleave(num_epsilon, dim=0)
|
||||
cond_eps = {
|
||||
key: cond[key].repeat_interleave(num_epsilon, dim=0) for key in cond
|
||||
}
|
||||
|
||||
def ode_func(t, x):
|
||||
x = x[:, :-1]
|
||||
@@ -132,7 +135,7 @@ def get_likelihood_fn(
|
||||
model_fn = model_ft
|
||||
else:
|
||||
model_fn = model
|
||||
x = x.view(shape) # B x horizon x transition_dim
|
||||
x = x.view(shape) # B x horizon x action_dim
|
||||
drift = drift_fn(
|
||||
model_fn,
|
||||
x,
|
||||
|
||||
@@ -25,6 +25,7 @@ class VisionDiffusionMLP(nn.Module):
|
||||
transition_dim,
|
||||
horizon_steps,
|
||||
cond_dim,
|
||||
img_cond_steps=1,
|
||||
time_dim=16,
|
||||
mlp_dims=[256, 256],
|
||||
activation_type="Mish",
|
||||
@@ -46,6 +47,8 @@ class VisionDiffusionMLP(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:
|
||||
@@ -101,36 +104,44 @@ class VisionDiffusionMLP(nn.Module):
|
||||
self,
|
||||
x,
|
||||
time,
|
||||
cond=None,
|
||||
cond: dict,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
x: (B,T,obs_dim)
|
||||
x: (B, Ta, Da)
|
||||
time: (B,) or int, diffusion step
|
||||
cond: dict (B,cond_step,cond_dim)
|
||||
output: (B,T,input_dim)
|
||||
cond: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
rgb: (B, To, C, H, W)
|
||||
|
||||
TODO long term: more flexible handling of cond
|
||||
"""
|
||||
# flatten T and input_dim
|
||||
B, T, input_dim = x.shape
|
||||
B, Ta, Da = x.shape
|
||||
_, T_rgb, C, H, W = cond["rgb"].shape
|
||||
|
||||
# flatten chunk
|
||||
x = x.view(B, -1)
|
||||
|
||||
# flatten cond_dim if exists
|
||||
if cond["rgb"].ndim == 5:
|
||||
rgb = einops.rearrange(cond["rgb"], "b d c h w -> (b d) c h w")
|
||||
# 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 = cond["rgb"]
|
||||
if cond["state"].ndim == 3:
|
||||
state = einops.rearrange(cond["state"], "b d c -> (b d) c")
|
||||
else:
|
||||
state = cond["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)
|
||||
@@ -141,7 +152,7 @@ class VisionDiffusionMLP(nn.Module):
|
||||
feat = torch.cat([feat1, feat2], dim=-1)
|
||||
else: # single image
|
||||
if self.augment:
|
||||
rgb = self.aug(rgb) # uint8 -> float32
|
||||
rgb = self.aug(rgb)
|
||||
feat = self.backbone(rgb)
|
||||
|
||||
# compress
|
||||
@@ -159,7 +170,7 @@ class VisionDiffusionMLP(nn.Module):
|
||||
|
||||
# mlp
|
||||
out = self.mlp_mean(x)
|
||||
return out.view(B, T, input_dim)
|
||||
return out.view(B, Ta, Da)
|
||||
|
||||
|
||||
class DiffusionMLP(nn.Module):
|
||||
@@ -210,27 +221,32 @@ class DiffusionMLP(nn.Module):
|
||||
self,
|
||||
x,
|
||||
time,
|
||||
cond=None,
|
||||
cond,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
x: (B,T,obs_dim)
|
||||
x: (B, Ta, Da)
|
||||
time: (B,) or int, diffusion step
|
||||
cond: (B,cond_step,cond_dim)
|
||||
output: (B,T,input_dim)
|
||||
cond: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
"""
|
||||
# flatten T and input_dim
|
||||
B, T, input_dim = x.shape
|
||||
B, Ta, Da = x.shape
|
||||
|
||||
# flatten chunk
|
||||
x = x.view(B, -1)
|
||||
cond = cond.view(B, -1) if cond is not None else None
|
||||
|
||||
# flatten history
|
||||
state = cond["state"].view(B, -1)
|
||||
|
||||
# obs encoder
|
||||
if hasattr(self, "cond_mlp"):
|
||||
cond = self.cond_mlp(cond)
|
||||
state = self.cond_mlp(state)
|
||||
|
||||
# append time and cond
|
||||
time = time.view(B, 1)
|
||||
time_emb = self.time_embedding(time).view(B, self.time_dim)
|
||||
x = torch.cat([x, time_emb, cond], dim=-1)
|
||||
x = torch.cat([x, time_emb, state], dim=-1)
|
||||
|
||||
# mlp
|
||||
# mlp head
|
||||
out = self.mlp_mean(x)
|
||||
return out.view(B, T, input_dim)
|
||||
return out.view(B, Ta, Da)
|
||||
|
||||
@@ -26,12 +26,6 @@ def extract(a, t, x_shape):
|
||||
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
|
||||
|
||||
|
||||
def apply_obs_conditioning(x, conditions, action_dim):
|
||||
for t, val in conditions.items():
|
||||
x[:, t, action_dim:] = val.clone()
|
||||
return x
|
||||
|
||||
|
||||
def make_timesteps(batch_size, i, device):
|
||||
t = torch.full((batch_size,), i, device=device, dtype=torch.long)
|
||||
return t
|
||||
|
||||
+13
-6
@@ -270,15 +270,22 @@ class Unet1D(nn.Module):
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
x: (B,T,input_dim)
|
||||
x: (B, Ta, act_dim)
|
||||
time: (B,) or int, diffusion step
|
||||
cond: (B,obs_step,cond_dim)
|
||||
output: (B,T,input_dim)
|
||||
cond: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, obs_dim)
|
||||
"""
|
||||
B = len(x)
|
||||
|
||||
# move chunk dim to the end
|
||||
x = einops.rearrange(x, "b h t -> b t h")
|
||||
cond = cond.view(cond.shape[0], -1)
|
||||
|
||||
# flatten history
|
||||
state = cond["state"].view(B, -1)
|
||||
|
||||
# obs encoder
|
||||
if hasattr(self, "cond_mlp"):
|
||||
cond = self.cond_mlp(cond)
|
||||
state = self.cond_mlp(state)
|
||||
|
||||
# 1. time
|
||||
if not torch.is_tensor(time):
|
||||
@@ -288,7 +295,7 @@ class Unet1D(nn.Module):
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
time = time.expand(x.shape[0])
|
||||
global_feature = self.time_mlp(time)
|
||||
global_feature = torch.cat([global_feature, cond], axis=-1)
|
||||
global_feature = torch.cat([global_feature, state], axis=-1)
|
||||
|
||||
# encode local features
|
||||
h_local = list()
|
||||
|
||||
@@ -1,6 +1,14 @@
|
||||
"""
|
||||
PPO for Gaussian policy.
|
||||
|
||||
To: observation sequence length
|
||||
Ta: action chunk size
|
||||
Do: observation dimension
|
||||
Da: action dimension
|
||||
|
||||
C: image channels
|
||||
H, W: image height and width
|
||||
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
@@ -41,8 +49,10 @@ class PPO_Gaussian(VPG_Gaussian):
|
||||
"""
|
||||
PPO loss
|
||||
|
||||
obs: (B, obs_step, obs_dim)
|
||||
actions: (B, horizon_step, action_dim)
|
||||
obs: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
rgb: (B, To, C, H, W)
|
||||
actions: (B, Ta, Da)
|
||||
returns: (B, )
|
||||
values: (B, )
|
||||
advantages: (B,)
|
||||
|
||||
@@ -29,9 +29,8 @@ class RWR_Gaussian(GaussianModel):
|
||||
|
||||
# override
|
||||
def loss(self, actions, obs, reward_weights):
|
||||
cond = obs
|
||||
B = cond.shape[0]
|
||||
means, scales = self.network(cond)
|
||||
B = len(obs)
|
||||
means, scales = self.network(obs)
|
||||
|
||||
dist = D.Normal(loc=means, scale=scales)
|
||||
log_prob = dist.log_prob(actions.view(B, -1)).mean(-1)
|
||||
@@ -42,16 +41,8 @@ class RWR_Gaussian(GaussianModel):
|
||||
# override
|
||||
@torch.no_grad()
|
||||
def forward(self, cond, deterministic=False, **kwargs):
|
||||
"""
|
||||
Args:
|
||||
cond: (batch_size, horizon, obs_dim)
|
||||
|
||||
Return:
|
||||
actions: (batch_size, horizon_steps, transition_dim)
|
||||
"""
|
||||
B = cond.shape[0]
|
||||
actions = super().forward(
|
||||
cond=cond.view(B, -1),
|
||||
cond=cond,
|
||||
deterministic=deterministic,
|
||||
randn_clip_value=self.randn_clip_value,
|
||||
)
|
||||
|
||||
+19
-34
@@ -15,26 +15,14 @@ class VPG_Gaussian(GaussianModel):
|
||||
self,
|
||||
actor,
|
||||
critic,
|
||||
cond_steps=1,
|
||||
randn_clip_value=10,
|
||||
network_path=None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(network=actor, **kwargs)
|
||||
self.cond_steps = cond_steps
|
||||
self.randn_clip_value = randn_clip_value
|
||||
|
||||
# Value function for obs - simple MLP
|
||||
self.critic = critic.to(self.device)
|
||||
if network_path is not None:
|
||||
checkpoint = torch.load(
|
||||
network_path, map_location=self.device, weights_only=True
|
||||
)
|
||||
self.load_state_dict(
|
||||
checkpoint["model"],
|
||||
strict=False,
|
||||
)
|
||||
logging.info("Loaded actor from %s", network_path)
|
||||
|
||||
# Re-name network to actor
|
||||
self.actor_ft = actor
|
||||
@@ -44,15 +32,31 @@ class VPG_Gaussian(GaussianModel):
|
||||
for param in self.actor.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
# ---------- Sampling ----------#
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
cond,
|
||||
deterministic=False,
|
||||
use_base_policy=False,
|
||||
):
|
||||
return super().forward(
|
||||
cond=cond,
|
||||
deterministic=deterministic,
|
||||
randn_clip_value=self.randn_clip_value,
|
||||
network_override=self.actor if use_base_policy else None,
|
||||
)
|
||||
|
||||
# ---------- RL training ----------#
|
||||
|
||||
def get_logprobs(
|
||||
self,
|
||||
cond,
|
||||
actions,
|
||||
use_base_policy=False,
|
||||
):
|
||||
B, T, D = actions.shape
|
||||
if not isinstance(cond, dict):
|
||||
cond = cond.view(B, -1)
|
||||
B = len(actions)
|
||||
dist = self.forward_train(
|
||||
cond,
|
||||
deterministic=False,
|
||||
@@ -66,22 +70,3 @@ class VPG_Gaussian(GaussianModel):
|
||||
|
||||
def loss(self, obs, actions, reward):
|
||||
raise NotImplementedError
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
cond,
|
||||
deterministic=False,
|
||||
use_base_policy=False,
|
||||
):
|
||||
if isinstance(cond, dict):
|
||||
B = cond["state"].shape[0]
|
||||
else:
|
||||
B = cond.shape[0]
|
||||
cond = cond.view(B, -1)
|
||||
return super().forward(
|
||||
cond=cond,
|
||||
deterministic=deterministic,
|
||||
randn_clip_value=self.randn_clip_value,
|
||||
network_override=self.actor if use_base_policy else None,
|
||||
)
|
||||
|
||||
+12
-2
@@ -1,6 +1,14 @@
|
||||
"""
|
||||
PPO for GMM policy.
|
||||
|
||||
To: observation sequence length
|
||||
Ta: action chunk size
|
||||
Do: observation dimension
|
||||
Da: action dimension
|
||||
|
||||
C: image channels
|
||||
H, W: image height and width
|
||||
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
@@ -41,8 +49,10 @@ class PPO_GMM(VPG_GMM):
|
||||
"""
|
||||
PPO loss
|
||||
|
||||
obs: (B, obs_step, obs_dim)
|
||||
actions: (B, horizon_step, action_dim)
|
||||
obs: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
rgb: (B, To, C, H, W)
|
||||
actions: (B, Ta, Da)
|
||||
returns: (B, )
|
||||
values: (B, )
|
||||
advantages: (B,)
|
||||
|
||||
+13
-23
@@ -8,36 +8,35 @@ class VPG_GMM(GMMModel):
|
||||
self,
|
||||
actor,
|
||||
critic,
|
||||
cond_steps=1,
|
||||
network_path=None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(network=actor, **kwargs)
|
||||
self.cond_steps = cond_steps
|
||||
|
||||
# Re-name network to actor
|
||||
self.actor_ft = actor
|
||||
|
||||
# Value function for obs - simple MLP
|
||||
self.critic = critic.to(self.device)
|
||||
if network_path is not None:
|
||||
checkpoint = torch.load(
|
||||
network_path, map_location=self.device, weights_only=True
|
||||
)
|
||||
self.load_state_dict(
|
||||
checkpoint["model"],
|
||||
strict=False,
|
||||
)
|
||||
logging.info("Loaded actor from %s", network_path)
|
||||
|
||||
# ---------- Sampling ----------#
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, cond, deterministic=False):
|
||||
return super().forward(
|
||||
cond=cond,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
|
||||
# ---------- RL training ----------#
|
||||
|
||||
def get_logprobs(
|
||||
self,
|
||||
cond,
|
||||
actions,
|
||||
):
|
||||
B, T, D = actions.shape
|
||||
B = len(actions)
|
||||
dist, entropy, std = self.forward_train(
|
||||
cond.view(B, -1),
|
||||
cond,
|
||||
deterministic=False,
|
||||
)
|
||||
log_prob = dist.log_prob(actions.view(B, -1))
|
||||
@@ -45,12 +44,3 @@ class VPG_GMM(GMMModel):
|
||||
|
||||
def loss(self, obs, chains, reward):
|
||||
raise NotImplementedError
|
||||
|
||||
# override to diffuse over action only
|
||||
@torch.no_grad()
|
||||
def forward(self, cond, deterministic=False):
|
||||
B = cond.shape[0]
|
||||
return super().forward(
|
||||
cond=cond.view(B, -1),
|
||||
deterministic=deterministic,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user