v0.5 to main (#10)
* v0.5 (#9) * update idql configs * update awr configs * update dipo configs * update qsm configs * update dqm configs * update project version to 0.5.0
This commit is contained in:
@@ -169,8 +169,11 @@ class DiffusionModel(nn.Module):
|
||||
|
||||
# ---------- Sampling ----------#
|
||||
|
||||
def p_mean_var(self, x, t, cond, index=None):
|
||||
noise = self.network(x, t, cond=cond)
|
||||
def p_mean_var(self, x, t, cond, index=None, network_override=None):
|
||||
if network_override is not None:
|
||||
noise = network_override(x, t, cond=cond)
|
||||
else:
|
||||
noise = self.network(x, t, cond=cond)
|
||||
|
||||
# Predict x_0
|
||||
if self.predict_epsilon:
|
||||
@@ -228,7 +231,7 @@ class DiffusionModel(nn.Module):
|
||||
return mu, logvar
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, cond):
|
||||
def forward(self, cond, deterministic=True):
|
||||
"""
|
||||
Forward pass for sampling actions. Used in evaluating pre-trained/fine-tuned policy. Not modifying diffusion clipping
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@ Actor and Critic models for model-free online RL with DIffusion POlicy (DIPO).
|
||||
|
||||
import torch
|
||||
import logging
|
||||
import copy
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
@@ -27,45 +28,67 @@ class DIPODiffusion(DiffusionModel):
|
||||
assert not self.use_ddim, "DQL does not support DDIM"
|
||||
self.critic = critic.to(self.device)
|
||||
|
||||
# target critic
|
||||
self.critic_target = copy.deepcopy(self.critic)
|
||||
|
||||
# reassign actor
|
||||
self.actor = self.network
|
||||
|
||||
# target actor
|
||||
self.actor_target = copy.deepcopy(self.actor)
|
||||
|
||||
# Minimum std used in denoising process when sampling action - helps exploration
|
||||
self.min_sampling_denoising_std = min_sampling_denoising_std
|
||||
|
||||
# ---------- RL training ----------#
|
||||
|
||||
def loss_critic(self, obs, next_obs, actions, rewards, dones, gamma):
|
||||
def loss_critic(self, obs, next_obs, actions, rewards, terminated, gamma):
|
||||
|
||||
# get current Q-function
|
||||
current_q1, current_q2 = self.critic(obs, actions)
|
||||
|
||||
# get next Q-function
|
||||
next_actions = self.forward(
|
||||
cond=next_obs,
|
||||
deterministic=False,
|
||||
) # 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)
|
||||
with torch.no_grad():
|
||||
# get next Q-function
|
||||
next_actions = self.forward(
|
||||
cond=next_obs,
|
||||
deterministic=False,
|
||||
) # forward() has no gradient, which is desired here.
|
||||
next_q1, next_q2 = self.critic_target(next_obs, next_actions)
|
||||
next_q = torch.min(next_q1, next_q2)
|
||||
|
||||
# terminal state mask
|
||||
mask = 1 - dones
|
||||
# terminal state mask
|
||||
mask = 1 - terminated
|
||||
|
||||
# flatten
|
||||
rewards = rewards.view(-1)
|
||||
next_q = next_q.view(-1)
|
||||
mask = mask.view(-1)
|
||||
# flatten
|
||||
rewards = rewards.view(-1)
|
||||
next_q = next_q.view(-1)
|
||||
mask = mask.view(-1)
|
||||
|
||||
# target value
|
||||
target_q = rewards + gamma * next_q * mask
|
||||
# target value
|
||||
target_q = rewards + gamma * next_q * mask
|
||||
|
||||
# Update critic
|
||||
loss_critic = torch.mean((current_q1 - target_q) ** 2) + torch.mean(
|
||||
(current_q2 - target_q) ** 2
|
||||
)
|
||||
|
||||
return loss_critic
|
||||
|
||||
def update_target_critic(self, tau):
|
||||
for target_param, source_param in zip(
|
||||
self.critic_target.parameters(), self.critic.parameters()
|
||||
):
|
||||
target_param.data.copy_(
|
||||
target_param.data * (1.0 - tau) + source_param.data * tau
|
||||
)
|
||||
|
||||
def update_target_actor(self, tau):
|
||||
for target_param, source_param in zip(
|
||||
self.actor_target.parameters(), self.actor.parameters()
|
||||
):
|
||||
target_param.data.copy_(
|
||||
target_param.data * (1.0 - tau) + source_param.data * tau
|
||||
)
|
||||
|
||||
# ---------- Sampling ----------#``
|
||||
|
||||
# override
|
||||
@@ -75,6 +98,7 @@ class DIPODiffusion(DiffusionModel):
|
||||
cond,
|
||||
deterministic=False,
|
||||
):
|
||||
"""Use target actor"""
|
||||
device = self.betas.device
|
||||
B = len(cond["state"])
|
||||
|
||||
@@ -87,6 +111,7 @@ class DIPODiffusion(DiffusionModel):
|
||||
x=x,
|
||||
t=t_b,
|
||||
cond=cond,
|
||||
network_override=self.actor_target,
|
||||
)
|
||||
std = torch.exp(0.5 * logvar)
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ Diffusion Q-Learning (DQL)
|
||||
import torch
|
||||
import logging
|
||||
import numpy as np
|
||||
import copy
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
@@ -28,6 +29,9 @@ class DQLDiffusion(DiffusionModel):
|
||||
assert not self.use_ddim, "DQL does not support DDIM"
|
||||
self.critic = critic.to(self.device)
|
||||
|
||||
# target critic
|
||||
self.critic_target = copy.deepcopy(self.critic)
|
||||
|
||||
# reassign actor
|
||||
self.actor = self.network
|
||||
|
||||
@@ -36,39 +40,46 @@ class DQLDiffusion(DiffusionModel):
|
||||
|
||||
# ---------- RL training ----------#
|
||||
|
||||
def loss_critic(self, obs, next_obs, actions, rewards, dones, gamma):
|
||||
def loss_critic(self, obs, next_obs, actions, rewards, terminated, gamma):
|
||||
|
||||
# get current Q-function
|
||||
current_q1, current_q2 = self.critic(obs, actions)
|
||||
|
||||
# get next Q-function
|
||||
next_actions = self.forward(
|
||||
cond=next_obs,
|
||||
deterministic=False,
|
||||
) # 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)
|
||||
with torch.no_grad():
|
||||
next_actions = self.forward(
|
||||
cond=next_obs,
|
||||
deterministic=False,
|
||||
) # forward() has no gradient, which is desired here.
|
||||
next_q1, next_q2 = self.critic_target(next_obs, next_actions)
|
||||
next_q = torch.min(next_q1, next_q2)
|
||||
|
||||
# terminal state mask
|
||||
mask = 1 - dones
|
||||
# terminal state mask
|
||||
mask = 1 - terminated
|
||||
|
||||
# flatten
|
||||
rewards = rewards.view(-1)
|
||||
next_q = next_q.view(-1)
|
||||
mask = mask.view(-1)
|
||||
# flatten
|
||||
rewards = rewards.view(-1)
|
||||
next_q = next_q.view(-1)
|
||||
mask = mask.view(-1)
|
||||
|
||||
# target value
|
||||
target_q = rewards + gamma * next_q * mask
|
||||
# target value
|
||||
target_q = rewards + gamma * next_q * mask
|
||||
|
||||
# Update critic
|
||||
loss_critic = torch.mean((current_q1 - target_q) ** 2) + torch.mean(
|
||||
(current_q2 - target_q) ** 2
|
||||
)
|
||||
|
||||
return loss_critic
|
||||
|
||||
def loss_actor(self, obs, actions, q1, q2, eta):
|
||||
bc_loss = self.loss(actions, obs)
|
||||
def loss_actor(self, obs, eta, act_steps):
|
||||
action_new = self.forward_train(
|
||||
cond=obs,
|
||||
deterministic=False,
|
||||
)[
|
||||
:, :act_steps
|
||||
] # with gradient
|
||||
q1, q2 = self.critic(obs, action_new)
|
||||
bc_loss = self.loss(action_new, obs)
|
||||
if np.random.uniform() > 0.5:
|
||||
q_loss = -q1.mean() / q2.abs().mean().detach()
|
||||
else:
|
||||
@@ -76,6 +87,14 @@ class DQLDiffusion(DiffusionModel):
|
||||
actor_loss = bc_loss + eta * q_loss
|
||||
return actor_loss
|
||||
|
||||
def update_target_critic(self, tau):
|
||||
for target_param, source_param in zip(
|
||||
self.critic_target.parameters(), self.critic.parameters()
|
||||
):
|
||||
target_param.data.copy_(
|
||||
target_param.data * (1.0 - tau) + source_param.data * tau
|
||||
)
|
||||
|
||||
# ---------- Sampling ----------#``
|
||||
|
||||
# override
|
||||
|
||||
@@ -20,11 +20,6 @@ def expectile_loss(diff, expectile=0.8):
|
||||
return weight * (diff**2)
|
||||
|
||||
|
||||
def soft_update(target, source, tau):
|
||||
for target_param, param in zip(target.parameters(), source.parameters()):
|
||||
target_param.data.copy_(target_param.data * (1.0 - tau) + param.data * tau)
|
||||
|
||||
|
||||
class IDQLDiffusion(RWRDiffusion):
|
||||
|
||||
def __init__(
|
||||
@@ -56,7 +51,6 @@ class IDQLDiffusion(RWRDiffusion):
|
||||
|
||||
# compute advantage
|
||||
adv = q - v
|
||||
|
||||
return adv
|
||||
|
||||
def loss_critic_v(self, obs, actions):
|
||||
@@ -64,10 +58,9 @@ class IDQLDiffusion(RWRDiffusion):
|
||||
|
||||
# get the value loss
|
||||
v_loss = expectile_loss(adv).mean()
|
||||
|
||||
return v_loss
|
||||
|
||||
def loss_critic_q(self, obs, next_obs, actions, rewards, dones, gamma):
|
||||
def loss_critic_q(self, obs, next_obs, actions, rewards, terminated, gamma):
|
||||
|
||||
# get current Q-function
|
||||
current_q1, current_q2 = self.critic_q(obs, actions)
|
||||
@@ -77,7 +70,7 @@ class IDQLDiffusion(RWRDiffusion):
|
||||
next_v = self.critic_v(next_obs)
|
||||
|
||||
# terminal state mask
|
||||
mask = 1 - dones
|
||||
mask = 1 - terminated
|
||||
|
||||
# flatten
|
||||
rewards = rewards.view(-1)
|
||||
@@ -91,11 +84,15 @@ class IDQLDiffusion(RWRDiffusion):
|
||||
q_loss = torch.mean((current_q1 - discounted_q) ** 2) + torch.mean(
|
||||
(current_q2 - discounted_q) ** 2
|
||||
)
|
||||
|
||||
return q_loss
|
||||
|
||||
def update_target_critic(self, tau):
|
||||
soft_update(self.target_q, self.critic_q, tau)
|
||||
for target_param, source_param in zip(
|
||||
self.target_q.parameters(), self.critic_q.parameters()
|
||||
):
|
||||
target_param.data.copy_(
|
||||
target_param.data * (1.0 - tau) + source_param.data * tau
|
||||
)
|
||||
|
||||
# override
|
||||
def p_losses(
|
||||
@@ -116,10 +113,9 @@ class IDQLDiffusion(RWRDiffusion):
|
||||
|
||||
# Loss with mask
|
||||
if self.predict_epsilon:
|
||||
loss = F.mse_loss(x_recon, noise, reduction="none")
|
||||
loss = F.mse_loss(x_recon, noise)
|
||||
else:
|
||||
loss = F.mse_loss(x_recon, x_start, reduction="none")
|
||||
loss = einops.reduce(loss, "b h d -> b", "mean")
|
||||
loss = F.mse_loss(x_recon, x_start)
|
||||
return loss.mean()
|
||||
|
||||
# ---------- Sampling ----------#``
|
||||
@@ -190,4 +186,4 @@ class IDQLDiffusion(RWRDiffusion):
|
||||
|
||||
# squeeze dummy dimension
|
||||
samples = samples_best[0]
|
||||
return samples
|
||||
return samples
|
||||
@@ -14,15 +14,18 @@ import torch
|
||||
import logging
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
from .diffusion_ppo import PPODiffusion
|
||||
from .diffusion_vpg import VPGDiffusion
|
||||
from .exact_likelihood import get_likelihood_fn
|
||||
|
||||
|
||||
class PPOExactDiffusion(PPODiffusion):
|
||||
class PPOExactDiffusion(VPGDiffusion):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
sde,
|
||||
clip_ploss_coef,
|
||||
clip_vloss_coef=None,
|
||||
norm_adv=True,
|
||||
sde_hutchinson_type="Rademacher",
|
||||
sde_rtol=1e-4,
|
||||
sde_atol=1e-4,
|
||||
@@ -41,6 +44,9 @@ class PPOExactDiffusion(PPODiffusion):
|
||||
self.betas,
|
||||
sde_min_beta,
|
||||
)
|
||||
self.clip_ploss_coef = clip_ploss_coef
|
||||
self.clip_vloss_coef = clip_vloss_coef
|
||||
self.norm_adv = norm_adv
|
||||
|
||||
# set up likelihood function
|
||||
self.likelihood_fn = get_likelihood_fn(
|
||||
@@ -62,7 +68,6 @@ class PPOExactDiffusion(PPODiffusion):
|
||||
|
||||
samples: (B x Ta x Da)
|
||||
"""
|
||||
# TODO: image input
|
||||
return self.likelihood_fn(
|
||||
self.actor,
|
||||
self.actor_ft,
|
||||
|
||||
@@ -14,16 +14,6 @@ log = logging.getLogger(__name__)
|
||||
from model.diffusion.diffusion_rwr import RWRDiffusion
|
||||
|
||||
|
||||
def expectile_loss(diff, expectile=0.8):
|
||||
weight = torch.where(diff > 0, expectile, (1 - expectile))
|
||||
return weight * (diff**2)
|
||||
|
||||
|
||||
def soft_update(target, source, tau):
|
||||
for target_param, param in zip(target.parameters(), source.parameters()):
|
||||
target_param.data.copy_(target_param.data * (1.0 - tau) + param.data * tau)
|
||||
|
||||
|
||||
class QSMDiffusion(RWRDiffusion):
|
||||
|
||||
def __init__(
|
||||
@@ -34,6 +24,8 @@ class QSMDiffusion(RWRDiffusion):
|
||||
):
|
||||
super().__init__(network=actor, **kwargs)
|
||||
self.critic_q = critic.to(self.device)
|
||||
|
||||
# target critic
|
||||
self.target_q = copy.deepcopy(critic)
|
||||
|
||||
# assign actor
|
||||
@@ -54,7 +46,6 @@ 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.requires_grad_(True)
|
||||
current_q1, current_q2 = self.critic_q(obs, x_noisy)
|
||||
|
||||
@@ -68,10 +59,10 @@ class QSMDiffusion(RWRDiffusion):
|
||||
|
||||
# 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()
|
||||
loss = F.mse_loss(-x_recon, q_grad_coeff * gradient_q)
|
||||
return loss
|
||||
|
||||
def loss_critic(self, obs, next_obs, actions, rewards, dones, gamma):
|
||||
def loss_critic(self, obs, next_obs, actions, rewards, terminated, gamma):
|
||||
|
||||
# get current Q-function
|
||||
current_q1, current_q2 = self.critic_q(obs, actions)
|
||||
@@ -86,7 +77,7 @@ class QSMDiffusion(RWRDiffusion):
|
||||
next_q = torch.min(next_q1, next_q2)
|
||||
|
||||
# terminal state mask
|
||||
mask = 1 - dones
|
||||
mask = 1 - terminated
|
||||
|
||||
# flatten
|
||||
rewards = rewards.view(-1)
|
||||
@@ -104,4 +95,9 @@ class QSMDiffusion(RWRDiffusion):
|
||||
return loss_critic
|
||||
|
||||
def update_target_critic(self, tau):
|
||||
soft_update(self.target_q, self.critic_q, tau)
|
||||
for target_param, source_param in zip(
|
||||
self.target_q.parameters(), self.critic_q.parameters()
|
||||
):
|
||||
target_param.data.copy_(
|
||||
target_param.data * (1.0 - tau) + source_param.data * tau
|
||||
)
|
||||
|
||||
@@ -298,7 +298,9 @@ class VPGDiffusion(DiffusionModel):
|
||||
|
||||
# clamp action at final step
|
||||
if self.final_action_clip_value is not None and i == len(t_all) - 1:
|
||||
x = torch.clamp(x, -self.final_action_clip_value, self.final_action_clip_value)
|
||||
x = torch.clamp(
|
||||
x, -self.final_action_clip_value, self.final_action_clip_value
|
||||
)
|
||||
|
||||
if return_chain:
|
||||
if not self.use_ddim and t <= self.ft_denoising_steps:
|
||||
|
||||
@@ -22,7 +22,7 @@ class VisionDiffusionMLP(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
backbone,
|
||||
transition_dim,
|
||||
action_dim,
|
||||
horizon_steps,
|
||||
cond_dim,
|
||||
img_cond_steps=1,
|
||||
@@ -77,9 +77,9 @@ class VisionDiffusionMLP(nn.Module):
|
||||
|
||||
# diffusion
|
||||
input_dim = (
|
||||
time_dim + transition_dim * horizon_steps + visual_feature_dim + cond_dim
|
||||
time_dim + action_dim * horizon_steps + visual_feature_dim + cond_dim
|
||||
)
|
||||
output_dim = transition_dim * horizon_steps
|
||||
output_dim = action_dim * horizon_steps
|
||||
self.time_embedding = nn.Sequential(
|
||||
SinusoidalPosEmb(time_dim),
|
||||
nn.Linear(time_dim, time_dim * 2),
|
||||
@@ -175,7 +175,7 @@ class DiffusionMLP(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transition_dim,
|
||||
action_dim,
|
||||
horizon_steps,
|
||||
cond_dim,
|
||||
time_dim=16,
|
||||
@@ -187,7 +187,7 @@ class DiffusionMLP(nn.Module):
|
||||
residual_style=False,
|
||||
):
|
||||
super().__init__()
|
||||
output_dim = transition_dim * horizon_steps
|
||||
output_dim = action_dim * horizon_steps
|
||||
self.time_embedding = nn.Sequential(
|
||||
SinusoidalPosEmb(time_dim),
|
||||
nn.Linear(time_dim, time_dim * 2),
|
||||
@@ -204,9 +204,9 @@ class DiffusionMLP(nn.Module):
|
||||
activation_type=activation_type,
|
||||
out_activation_type="Identity",
|
||||
)
|
||||
input_dim = time_dim + transition_dim * horizon_steps + cond_mlp_dims[-1]
|
||||
input_dim = time_dim + action_dim * horizon_steps + cond_mlp_dims[-1]
|
||||
else:
|
||||
input_dim = time_dim + transition_dim * horizon_steps + cond_dim
|
||||
input_dim = time_dim + action_dim * horizon_steps + cond_dim
|
||||
self.mlp_mean = model(
|
||||
[input_dim] + mlp_dims + [output_dim],
|
||||
activation_type=activation_type,
|
||||
|
||||
@@ -120,7 +120,7 @@ class Unet1D(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transition_dim,
|
||||
action_dim,
|
||||
cond_dim=None,
|
||||
diffusion_step_embed_dim=32,
|
||||
dim=32,
|
||||
@@ -134,7 +134,7 @@ class Unet1D(nn.Module):
|
||||
groupnorm_eps=1e-5,
|
||||
):
|
||||
super().__init__()
|
||||
dims = [transition_dim, *map(lambda m: dim * m, dim_mults)]
|
||||
dims = [action_dim, *map(lambda m: dim * m, dim_mults)]
|
||||
in_out = list(zip(dims[:-1], dims[1:]))
|
||||
log.info(f"Channel dimensions: {in_out}")
|
||||
|
||||
@@ -259,7 +259,7 @@ class Unet1D(nn.Module):
|
||||
activation_type=activation_type,
|
||||
eps=groupnorm_eps,
|
||||
),
|
||||
nn.Conv1d(dim, transition_dim, 1),
|
||||
nn.Conv1d(dim, action_dim, 1),
|
||||
)
|
||||
|
||||
def forward(
|
||||
|
||||
Reference in New Issue
Block a user