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())
+8 -59
View File
@@ -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)
+9 -12
View File
@@ -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)
+12 -20
View File
@@ -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)
+43 -48
View File
@@ -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()
+15 -4
View File
@@ -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
+18 -3
View File
@@ -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)
+12 -15
View File
@@ -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
+4 -16
View File
@@ -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)
+65 -82
View File
@@ -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
View File
@@ -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)
+8 -5
View File
@@ -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,
+46 -30
View File
@@ -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)
-6
View File
@@ -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
View File
@@ -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()
+12 -2
View File
@@ -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,)
+3 -12
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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,
)