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:
Allen Z. Ren
2024-10-07 16:35:13 -04:00
committed by GitHub
parent dd14c5887c
commit e0842e71dc
267 changed files with 6769 additions and 1645 deletions
+26 -26
View File
@@ -5,7 +5,6 @@ Critic networks.
from typing import Union
import torch
import copy
import einops
from copy import deepcopy
@@ -28,20 +27,15 @@ class CriticObs(torch.nn.Module):
super().__init__()
mlp_dims = [cond_dim] + mlp_dims + [1]
if residual_style:
self.Q1 = ResidualMLP(
mlp_dims,
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
)
model = ResidualMLP
else:
self.Q1 = MLP(
mlp_dims,
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
verbose=False,
)
model = MLP
self.Q1 = model(
mlp_dims,
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
)
def forward(self, cond: Union[dict, torch.Tensor]):
"""
@@ -72,26 +66,28 @@ class CriticObsAct(torch.nn.Module):
activation_type="Mish",
use_layernorm=False,
residual_tyle=False,
double_q=True,
**kwargs,
):
super().__init__()
mlp_dims = [cond_dim + action_dim * action_steps] + mlp_dims + [1]
if residual_tyle:
self.Q1 = ResidualMLP(
mlp_dims,
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
)
model = ResidualMLP
else:
self.Q1 = MLP(
model = MLP
self.Q1 = model(
mlp_dims,
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
)
if double_q:
self.Q2 = model(
mlp_dims,
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
verbose=False,
)
self.Q2 = copy.deepcopy(self.Q1)
def forward(self, cond: dict, action):
"""
@@ -108,9 +104,13 @@ class CriticObsAct(torch.nn.Module):
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)
if hasattr(self, "Q2"):
q1 = self.Q1(x)
q2 = self.Q2(x)
return q1.squeeze(1), q2.squeeze(1)
else:
q1 = self.Q1(x)
return q1.squeeze(1)
class ViTCritic(CriticObs):
+27 -3
View File
@@ -19,13 +19,16 @@ class GaussianModel(torch.nn.Module):
network_path=None,
device="cuda:0",
randn_clip_value=10,
tanh_output=False,
):
super().__init__()
self.device = device
self.network = network.to(device)
if network_path is not None:
checkpoint = torch.load(
network_path, map_location=self.device, weights_only=True
network_path,
map_location=self.device,
weights_only=True,
)
self.load_state_dict(
checkpoint["model"],
@@ -40,12 +43,16 @@ class GaussianModel(torch.nn.Module):
# Clip sampled randn (from standard deviation) such that the sampled action is not too far away from mean
self.randn_clip_value = randn_clip_value
# Whether to apply tanh to the **sampled** action --- used in SAC
self.tanh_output = tanh_output
def loss(
self,
true_action,
cond,
ent_coef,
):
"""no squashing"""
B = len(true_action)
dist = self.forward_train(
cond,
@@ -80,6 +87,8 @@ class GaussianModel(torch.nn.Module):
cond,
deterministic=False,
network_override=None,
reparameterize=False,
get_logprob=False,
):
B = len(cond["state"]) if "state" in cond else len(cond["rgb"])
T = self.horizon_steps
@@ -88,9 +97,24 @@ class GaussianModel(torch.nn.Module):
deterministic=deterministic,
network_override=network_override,
)
sampled_action = dist.sample()
if reparameterize:
sampled_action = dist.rsample()
else:
sampled_action = dist.sample()
sampled_action.clamp_(
dist.loc - self.randn_clip_value * dist.scale,
dist.loc + self.randn_clip_value * dist.scale,
)
return sampled_action.view(B, T, -1)
if get_logprob:
log_prob = dist.log_prob(sampled_action)
# For SAC/RLPD, squash mean after sampling here instead of right after model output as in PPO
if self.tanh_output:
sampled_action = torch.tanh(sampled_action)
log_prob -= torch.log(1 - sampled_action.pow(2) + 1e-6)
return sampled_action.view(B, T, -1), log_prob.sum(1, keepdim=False)
else:
if self.tanh_output:
sampled_action = torch.tanh(sampled_action)
return sampled_action.view(B, T, -1)
+24 -35
View File
@@ -7,7 +7,6 @@ Residual model is taken from https://github.com/ALRhub/d3il/blob/main/agents/mod
import torch
from torch import nn
from torch.nn.utils import spectral_norm
from collections import OrderedDict
import logging
@@ -26,7 +25,6 @@ activation_dict = nn.ModuleDict(
class MLP(nn.Module):
def __init__(
self,
dim_list,
@@ -35,7 +33,9 @@ class MLP(nn.Module):
activation_type="Tanh",
out_activation_type="Identity",
use_layernorm=False,
use_spectralnorm=False,
use_layernorm_final=False,
dropout=0,
use_drop_final=False,
verbose=False,
):
super(MLP, self).__init__()
@@ -50,39 +50,25 @@ class MLP(nn.Module):
o_dim = dim_list[idx + 1]
if append_dim > 0 and idx in append_layers:
i_dim += append_dim
linear_layer = nn.Linear(i_dim, o_dim)
if use_spectralnorm:
linear_layer = spectral_norm(linear_layer)
if idx == num_layer - 1:
module = nn.Sequential(
OrderedDict(
[
("linear_1", linear_layer),
("act_1", activation_dict[out_activation_type]),
]
)
)
else:
if use_layernorm:
module = nn.Sequential(
OrderedDict(
[
("linear_1", linear_layer),
("norm_1", nn.LayerNorm(o_dim)),
("act_1", activation_dict[activation_type]),
]
)
)
else:
module = nn.Sequential(
OrderedDict(
[
("linear_1", linear_layer),
("act_1", activation_dict[activation_type]),
]
)
)
# Add module components
layers = [("linear_1", linear_layer)]
if use_layernorm and (idx < num_layer - 1 or use_layernorm_final):
layers.append(("norm_1", nn.LayerNorm(o_dim)))
if dropout > 0 and (idx < num_layer - 1 or use_drop_final):
layers.append(("dropout_1", nn.Dropout(dropout)))
# add activation function
act = (
activation_dict[activation_type]
if idx != num_layer - 1
else activation_dict[out_activation_type]
)
layers.append(("act_1", act))
# re-construct module
module = nn.Sequential(OrderedDict(layers))
self.moduleList.append(module)
if verbose:
logging.info(self.moduleList)
@@ -109,6 +95,7 @@ class ResidualMLP(nn.Module):
activation_type="Mish",
out_activation_type="Identity",
use_layernorm=False,
use_layernorm_final=False,
):
super(ResidualMLP, self).__init__()
hidden_dim = dim_list[1]
@@ -126,6 +113,8 @@ class ResidualMLP(nn.Module):
]
)
self.layers.append(nn.Linear(hidden_dim, dim_list[-1]))
if use_layernorm_final:
self.layers.append(nn.LayerNorm(dim_list[-1]))
self.layers.append(activation_dict[out_activation_type])
def forward(self, x):
+52 -29
View File
@@ -18,7 +18,7 @@ class Gaussian_VisionMLP(nn.Module):
def __init__(
self,
backbone,
transition_dim,
action_dim,
horizon_steps,
cond_dim,
img_cond_steps=1,
@@ -74,10 +74,10 @@ class Gaussian_VisionMLP(nn.Module):
)
# head
self.transition_dim = transition_dim
self.action_dim = action_dim
self.horizon_steps = horizon_steps
input_dim = visual_feature_dim + cond_dim
output_dim = transition_dim * horizon_steps
output_dim = action_dim * horizon_steps
if residual_style:
model = ResidualMLP
else:
@@ -97,7 +97,7 @@ class Gaussian_VisionMLP(nn.Module):
)
elif learn_fixed_std: # initialize to fixed_std
self.logvar = torch.nn.Parameter(
torch.log(torch.tensor([fixed_std**2 for _ in range(transition_dim)])),
torch.log(torch.tensor([fixed_std**2 for _ in range(action_dim)])),
requires_grad=True,
)
self.logvar_min = torch.nn.Parameter(
@@ -159,19 +159,19 @@ class Gaussian_VisionMLP(nn.Module):
x_encoded = torch.cat([feat, state], dim=-1)
out_mean = self.mlp_mean(x_encoded)
out_mean = torch.tanh(out_mean).view(
B, self.horizon_steps * self.transition_dim
B, self.horizon_steps * self.action_dim
) # tanh squashing in [-1, 1]
if self.learn_fixed_std:
out_logvar = torch.clamp(self.logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
out_scale = out_scale.view(1, self.transition_dim)
out_scale = out_scale.view(1, self.action_dim)
out_scale = out_scale.repeat(B, self.horizon_steps)
elif self.use_fixed_std:
out_scale = torch.ones_like(out_mean).to(device) * self.fixed_std
else:
out_logvar = self.mlp_logvar(x_encoded).view(
B, self.horizon_steps * self.transition_dim
B, self.horizon_steps * self.action_dim
)
out_logvar = torch.clamp(out_logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
@@ -179,48 +179,65 @@ class Gaussian_VisionMLP(nn.Module):
class Gaussian_MLP(nn.Module):
def __init__(
self,
transition_dim,
action_dim,
horizon_steps,
cond_dim,
mlp_dims=[256, 256, 256],
activation_type="Mish",
tanh_output=True, # sometimes we want to apply tanh after sampling instead of here, e.g., in SAC
residual_style=False,
use_layernorm=False,
dropout=0.0,
fixed_std=None,
learn_fixed_std=False,
std_min=0.01,
std_max=1,
):
super().__init__()
self.transition_dim = transition_dim
self.action_dim = action_dim
self.horizon_steps = horizon_steps
input_dim = cond_dim
output_dim = transition_dim * horizon_steps
output_dim = action_dim * horizon_steps
if residual_style:
model = ResidualMLP
else:
model = MLP
self.mlp_mean = model(
[input_dim] + mlp_dims + [output_dim],
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
)
if fixed_std is None:
# learning std
self.mlp_base = model(
[input_dim] + mlp_dims,
activation_type=activation_type,
out_activation_type=activation_type,
use_layernorm=use_layernorm,
use_layernorm_final=use_layernorm,
)
self.mlp_mean = MLP(
mlp_dims[-1:] + [output_dim],
out_activation_type="Identity",
)
self.mlp_logvar = MLP(
[input_dim] + mlp_dims[-1:] + [output_dim],
mlp_dims[-1:] + [output_dim],
out_activation_type="Identity",
)
else:
# no separate head for mean and std
self.mlp_mean = model(
[input_dim] + mlp_dims + [output_dim],
activation_type=activation_type,
out_activation_type="Identity",
use_layernorm=use_layernorm,
dropout=dropout,
)
elif learn_fixed_std: # initialize to fixed_std
self.logvar = torch.nn.Parameter(
torch.log(torch.tensor([fixed_std**2 for _ in range(transition_dim)])),
requires_grad=True,
)
if learn_fixed_std:
# initialize to fixed_std
self.logvar = torch.nn.Parameter(
torch.log(
torch.tensor([fixed_std**2 for _ in range(action_dim)])
),
requires_grad=True,
)
self.logvar_min = torch.nn.Parameter(
torch.log(torch.tensor(std_min**2)), requires_grad=False
)
@@ -230,6 +247,7 @@ class Gaussian_MLP(nn.Module):
self.use_fixed_std = fixed_std is not None
self.fixed_std = fixed_std
self.learn_fixed_std = learn_fixed_std
self.tanh_output = tanh_output
def forward(self, cond):
B = len(cond["state"])
@@ -239,22 +257,27 @@ class Gaussian_MLP(nn.Module):
state = cond["state"].view(B, -1)
# mlp
if hasattr(self, "mlp_base"):
state = self.mlp_base(state)
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]
if self.tanh_output:
out_mean = torch.tanh(out_mean)
out_mean = out_mean.view(B, self.horizon_steps * self.action_dim)
if self.learn_fixed_std:
out_logvar = torch.clamp(self.logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
out_scale = out_scale.view(1, self.transition_dim)
out_scale = out_scale.view(1, self.action_dim)
out_scale = out_scale.repeat(B, self.horizon_steps)
elif self.use_fixed_std:
out_scale = torch.ones_like(out_mean).to(device) * self.fixed_std
else:
out_logvar = self.mlp_logvar(state).view(
B, self.horizon_steps * self.transition_dim
B, self.horizon_steps * self.action_dim
)
out_logvar = torch.clamp(out_logvar, self.logvar_min, self.logvar_max)
out_logvar = torch.tanh(out_logvar)
out_logvar = self.logvar_min + 0.5 * (self.logvar_max - self.logvar_min) * (
out_logvar + 1
) # put back to full range
out_scale = torch.exp(0.5 * out_logvar)
return out_mean, out_scale
+7 -7
View File
@@ -12,7 +12,7 @@ class GMM_MLP(nn.Module):
def __init__(
self,
transition_dim,
action_dim,
horizon_steps,
cond_dim=None,
mlp_dims=[256, 256, 256],
@@ -26,10 +26,10 @@ class GMM_MLP(nn.Module):
std_max=1,
):
super().__init__()
self.transition_dim = transition_dim
self.action_dim = action_dim
self.horizon_steps = horizon_steps
input_dim = cond_dim
output_dim = transition_dim * horizon_steps * num_modes
output_dim = action_dim * horizon_steps * num_modes
self.num_modes = num_modes
if residual_style:
model = ResidualMLP
@@ -54,7 +54,7 @@ class GMM_MLP(nn.Module):
self.logvar = torch.nn.Parameter(
torch.log(
torch.tensor(
[fixed_std**2 for _ in range(transition_dim * num_modes)]
[fixed_std**2 for _ in range(action_dim * num_modes)]
)
),
requires_grad=True,
@@ -87,19 +87,19 @@ class GMM_MLP(nn.Module):
# mlp
out_mean = self.mlp_mean(state)
out_mean = torch.tanh(out_mean).view(
B, self.num_modes, self.horizon_steps * self.transition_dim
B, self.num_modes, self.horizon_steps * self.action_dim
) # tanh squashing in [-1, 1]
if self.learn_fixed_std:
out_logvar = torch.clamp(self.logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
out_scale = out_scale.view(1, self.num_modes, self.transition_dim)
out_scale = out_scale.view(1, self.num_modes, self.action_dim)
out_scale = out_scale.repeat(B, 1, self.horizon_steps)
elif self.use_fixed_std:
out_scale = torch.ones_like(out_mean).to(device) * self.fixed_std
else:
out_logvar = self.mlp_logvar(state).view(
B, self.num_modes, self.horizon_steps * self.transition_dim
B, self.num_modes, self.horizon_steps * self.action_dim
)
out_logvar = torch.clamp(out_logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
+21 -22
View File
@@ -16,7 +16,7 @@ logger = logging.getLogger(__name__)
class Gaussian_Transformer(nn.Module):
def __init__(
self,
transition_dim,
action_dim,
horizon_steps,
cond_dim,
transformer_embed_dim=256,
@@ -32,16 +32,16 @@ class Gaussian_Transformer(nn.Module):
):
super().__init__()
self.transition_dim = transition_dim
self.action_dim = action_dim
self.horizon_steps = horizon_steps
output_dim = transition_dim
output_dim = action_dim
if fixed_std is None: # learn the logvar
output_dim *= 2 # mean and logvar
logger.info("Using learned std")
elif learn_fixed_std: # learn logvar
self.logvar = torch.nn.Parameter(
torch.log(torch.tensor([fixed_std**2 for _ in range(transition_dim)])),
torch.log(torch.tensor([fixed_std**2 for _ in range(action_dim)])),
requires_grad=True,
)
logger.info(f"Using fixed std {fixed_std} with learning")
@@ -81,19 +81,19 @@ class Gaussian_Transformer(nn.Module):
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])
out_mean = out_mean.view(B, self.horizon_steps * self.transition_dim)
out_mean = torch.tanh(out[:, :, : self.action_dim])
out_mean = out_mean.view(B, self.horizon_steps * self.action_dim)
if self.learn_fixed_std:
out_logvar = torch.clamp(self.logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
out_scale = out_scale.view(1, self.transition_dim)
out_scale = out_scale.view(1, self.action_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(device) * self.fixed_std
else:
out_logvar = out[:, :, self.transition_dim :]
out_logvar = out_logvar.reshape(B, self.horizon_steps * self.transition_dim)
out_logvar = out[:, :, self.action_dim :]
out_logvar = out_logvar.reshape(B, self.horizon_steps * self.action_dim)
out_logvar = torch.clamp(out_logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
return out_mean, out_scale
@@ -102,7 +102,7 @@ class Gaussian_Transformer(nn.Module):
class GMM_Transformer(nn.Module):
def __init__(
self,
transition_dim,
action_dim,
horizon_steps,
cond_dim,
num_modes=5,
@@ -120,13 +120,12 @@ class GMM_Transformer(nn.Module):
super().__init__()
self.num_modes = num_modes
self.transition_dim = transition_dim
self.action_dim = action_dim
self.horizon_steps = horizon_steps
output_dim = transition_dim * num_modes
# + num_modes # mean and modes
output_dim = action_dim * num_modes
if fixed_std is None:
output_dim += num_modes * transition_dim # logvar for each mode
output_dim += num_modes * action_dim # logvar for each mode
logger.info("Using learned std")
elif (
learn_fixed_std
@@ -134,7 +133,7 @@ class GMM_Transformer(nn.Module):
self.logvar = torch.nn.Parameter(
torch.log(
torch.tensor(
[fixed_std**2 for _ in range(num_modes * transition_dim)]
[fixed_std**2 for _ in range(num_modes * action_dim)]
)
),
requires_grad=True,
@@ -179,32 +178,32 @@ class GMM_Transformer(nn.Module):
) # (B,horizon,output_dim), (B,horizon,emb_dim)
# use the first half of the output as mean
out_mean = torch.tanh(out[:, :, : self.num_modes * self.transition_dim])
out_mean = torch.tanh(out[:, :, : self.num_modes * self.action_dim])
out_mean = out_mean.reshape(
B, self.horizon_steps, self.num_modes, self.transition_dim
B, self.horizon_steps, self.num_modes, self.action_dim
)
out_mean = out_mean.permute(0, 2, 1, 3) # flip horizons and modes
out_mean = out_mean.reshape(
B, self.num_modes, self.horizon_steps * self.transition_dim
B, self.num_modes, self.horizon_steps * self.action_dim
)
if self.learn_fixed_std:
out_logvar = torch.clamp(self.logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
out_scale = out_scale.view(1, self.num_modes, self.transition_dim)
out_scale = out_scale.view(1, self.num_modes, self.action_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(device) * self.fixed_std
else:
out_logvar = out[
:, :, self.num_modes * self.transition_dim : -self.num_modes
:, :, self.num_modes * self.action_dim : -self.num_modes
]
out_logvar = out_logvar.reshape(
B, self.horizon_steps, self.num_modes, self.transition_dim
B, self.horizon_steps, self.num_modes, self.action_dim
)
out_logvar = out_logvar.permute(0, 2, 1, 3) # flip horizons and modes
out_logvar = out_logvar.reshape(
B, self.num_modes, self.horizon_steps * self.transition_dim
B, self.num_modes, self.horizon_steps * self.action_dim
)
out_logvar = torch.clamp(out_logvar, self.logvar_min, self.logvar_max)
out_scale = torch.exp(0.5 * out_logvar)
+6 -3
View File
@@ -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
+42 -17
View File
@@ -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)
+37 -18
View File
@@ -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
+11 -15
View File
@@ -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
+8 -3
View File
@@ -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,
+11 -15
View File
@@ -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
)
+3 -1
View File
@@ -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:
+7 -7
View File
@@ -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,
+3 -3
View File
@@ -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(
+187
View File
@@ -0,0 +1,187 @@
"""
Calibrated Conservative Q-Learning (CalQL) for Gaussian policy.
"""
import torch
import torch.nn as nn
import logging
from copy import deepcopy
import numpy as np
import einops
from model.common.gaussian import GaussianModel
log = logging.getLogger(__name__)
class CalQL_Gaussian(GaussianModel):
def __init__(
self,
actor,
critic,
network_path=None,
cql_clip_diff_min=-np.inf,
cql_clip_diff_max=np.inf,
cql_min_q_weight=5.0,
cql_n_actions=10,
**kwargs,
):
super().__init__(network=actor, network_path=None, **kwargs)
self.cql_clip_diff_min = cql_clip_diff_min
self.cql_clip_diff_max = cql_clip_diff_max
self.cql_min_q_weight = cql_min_q_weight
self.cql_n_actions = cql_n_actions
# initialize critic networks
self.critic = critic.to(self.device)
self.target_critic = deepcopy(critic).to(self.device)
# Load pre-trained checkpoint - note we are also loading the pre-trained critic here
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=True,
)
log.info("Loaded actor from %s", network_path)
log.info(
f"Number of network parameters: {sum(p.numel() for p in self.parameters())}"
)
def loss_critic(
self,
obs,
next_obs,
actions,
random_actions,
rewards,
returns,
terminated,
gamma,
alpha,
):
B = len(actions)
# Get initial TD loss
q_data1, q_data2 = self.critic(obs, actions)
with torch.no_grad():
# repeat for action samples
next_obs["state"] = next_obs["state"].repeat_interleave(
self.cql_n_actions, dim=0
)
# Get the next actions and logprobs
next_actions, next_logprobs = self.forward(
next_obs,
deterministic=False,
get_logprob=True,
)
next_q1, next_q2 = self.target_critic(next_obs, next_actions)
next_q = torch.min(next_q1, next_q2)
# Reshape the next_q to match the number of samples
next_q = next_q.view(B, self.cql_n_actions) # (B, n_sample)
next_logprobs = next_logprobs.view(B, self.cql_n_actions) # (B, n_sample)
# Get the max indices over the samples, and index into the next_q and next_log_probs
max_idx = torch.argmax(next_q, dim=1)
next_q = next_q[torch.arange(B), max_idx]
next_logprobs = next_logprobs[torch.arange(B), max_idx]
# Get the target Q values
target_q = rewards + gamma * (1 - terminated) * next_q
# Subtract the entropy bonus
target_q = target_q - alpha * next_logprobs
# TD loss
td_loss_1 = nn.functional.mse_loss(q_data1, target_q)
td_loss_2 = nn.functional.mse_loss(q_data2, target_q)
# Get actions and logprobs
log_rand_pi = 0.5 ** torch.prod(torch.tensor(random_actions.shape[-2:]))
pi_actions, log_pi = self.forward(
obs,
deterministic=False,
reparameterize=False,
get_logprob=True,
) # no gradient
# Random action Q values
n_random_actions = random_actions.shape[1]
obs_sample_state = {
"state": obs["state"].repeat_interleave(n_random_actions, dim=0)
}
random_actions = einops.rearrange(random_actions, "B N H A -> (B N) H A")
# Get the random action Q-values
q_rand_1, q_rand_2 = self.critic(obs_sample_state, random_actions)
q_rand_1 = q_rand_1 - log_rand_pi
q_rand_2 = q_rand_2 - log_rand_pi
# Reshape the random action Q values to match the number of samples
q_rand_1 = q_rand_1.view(B, n_random_actions) # (n_sample, B)
q_rand_2 = q_rand_2.view(B, n_random_actions)
# Policy action Q values
q_pi_1, q_pi_2 = self.critic(obs, pi_actions)
q_pi_1 = q_pi_1 - log_pi
q_pi_2 = q_pi_2 - log_pi
# Ensure calibration w.r.t. value function estimate
q_pi_1 = torch.max(q_pi_1, returns)[:, None] # (B, 1)
q_pi_2 = torch.max(q_pi_2, returns)[:, None] # (B, 1)
cat_q_1 = torch.cat([q_rand_1, q_pi_1], dim=-1) # (B, num_samples+1)
cql_qf1_ood = torch.logsumexp(cat_q_1, dim=-1) # max over num_samples
cat_q_2 = torch.cat([q_rand_2, q_pi_2], dim=-1) # (B, num_samples+1)
cql_qf2_ood = torch.logsumexp(cat_q_2, dim=-1) # sum over num_samples
# Subtract the log likelihood of the data
cql_qf1_diff = torch.clamp(
cql_qf1_ood - q_data1,
min=self.cql_clip_diff_min,
max=self.cql_clip_diff_max,
).mean()
cql_qf2_diff = torch.clamp(
cql_qf2_ood - q_data2,
min=self.cql_clip_diff_min,
max=self.cql_clip_diff_max,
).mean()
cql_min_qf1_loss = cql_qf1_diff * self.cql_min_q_weight
cql_min_qf2_loss = cql_qf2_diff * self.cql_min_q_weight
# Sum the two losses
critic_loss = td_loss_1 + td_loss_2 + cql_min_qf1_loss + cql_min_qf2_loss
return critic_loss
def loss_actor(self, obs, alpha):
action, logprob = self.forward(
obs,
deterministic=False,
reparameterize=True,
get_logprob=True,
)
q1, q2 = self.critic(obs, action)
actor_loss = -torch.min(q1, q2) + alpha * logprob
return actor_loss.mean()
def loss_temperature(self, obs, alpha, target_entropy):
with torch.no_grad():
_, logprob = self.forward(
obs,
deterministic=False,
get_logprob=True,
)
loss_alpha = -torch.mean(alpha * (logprob + target_entropy))
return loss_alpha
def update_target_critic(self, tau):
for target_param, param in zip(
self.target_critic.parameters(), self.critic.parameters()
):
target_param.data.copy_(tau * param.data + (1 - tau) * target_param.data)
+205
View File
@@ -0,0 +1,205 @@
"""
Imitation Bootstrapped Reinforcement Learning (IBRL) for Gaussian policy.
"""
import torch
import torch.nn as nn
import logging
from copy import deepcopy
from model.common.gaussian import GaussianModel
log = logging.getLogger(__name__)
class IBRL_Gaussian(GaussianModel):
def __init__(
self,
actor,
critic,
n_critics,
soft_action_sample=False,
soft_action_sample_beta=0.1,
**kwargs,
):
super().__init__(network=actor, **kwargs)
self.soft_action_sample = soft_action_sample
self.soft_action_sample_beta = soft_action_sample_beta
# Set up target actor
self.target_actor = deepcopy(actor)
# Frozen pre-trained policy
self.bc_policy = deepcopy(actor)
for param in self.bc_policy.parameters():
param.requires_grad = False
# initialize critic networks
self.critic_networks = [
deepcopy(critic).to(self.device) for _ in range(n_critics)
]
self.critic_networks = nn.ModuleList(self.critic_networks)
# initialize target networks
self.target_networks = [
deepcopy(critic).to(self.device) for _ in range(n_critics)
]
self.target_networks = nn.ModuleList(self.target_networks)
# Construct a "stateless" version of one of the models. It is "stateless" in the sense that the parameters are meta Tensors and do not have storage.
base_model = deepcopy(self.critic_networks[0])
self.base_model = base_model.to("meta")
self.ensemble_params, self.ensemble_buffers = torch.func.stack_module_state(
self.critic_networks
)
def critic_wrapper(self, params, buffers, data):
"""for vmap"""
return torch.func.functional_call(self.base_model, (params, buffers), data)
def get_random_indices(self, sz=None, num_ind=2):
"""get num_ind random indices from a set of size sz (used for getting critic targets)"""
if sz is None:
sz = len(self.critic_networks)
perm = torch.randperm(sz)
ind = perm[:num_ind].to(self.device)
return ind
def loss_critic(
self,
obs,
next_obs,
actions,
rewards,
terminated,
gamma,
):
# get random critic index
q1_ind, q2_ind = self.get_random_indices()
with torch.no_grad():
next_actions_bc = super().forward(
cond=next_obs,
deterministic=True,
network_override=self.bc_policy,
)
next_actions_rl = super().forward(
cond=next_obs,
deterministic=False,
network_override=self.target_actor,
)
# get the BC Q value
next_q1_bc = self.target_networks[q1_ind](next_obs, next_actions_bc)
next_q2_bc = self.target_networks[q2_ind](next_obs, next_actions_bc)
next_q_bc = torch.min(next_q1_bc, next_q2_bc)
# get the RL Q value
next_q1_rl = self.target_networks[q1_ind](next_obs, next_actions_rl)
next_q2_rl = self.target_networks[q2_ind](next_obs, next_actions_rl)
next_q_rl = torch.min(next_q1_rl, next_q2_rl)
# take the max Q value
next_q = torch.where(next_q_bc > next_q_rl, next_q_bc, next_q_rl)
# target value
target_q = rewards + gamma * (1 - terminated) * next_q # (B,)
# run all critics in batch
current_q = torch.vmap(self.critic_wrapper, in_dims=(0, 0, None))(
self.ensemble_params, self.ensemble_buffers, (obs, actions)
) # (n_critics, B)
loss_critic = torch.mean((current_q - target_q[None]) ** 2)
return loss_critic
def loss_actor(self, obs):
action = super().forward(
obs,
deterministic=False,
reparameterize=True,
) # use online policy only, also IBRL does not use tanh squashing
current_q = torch.vmap(self.critic_wrapper, in_dims=(0, 0, None))(
self.ensemble_params, self.ensemble_buffers, (obs, action)
) # (n_critics, B)
current_q = current_q.min(
dim=0
).values # unlike RLPD, IBRL uses the min Q value for actor update
loss_actor = -torch.mean(current_q)
return loss_actor
def update_target_critic(self, tau):
"""need to use ensemble_params instead of critic_networks"""
for target_ind, target_critic in enumerate(self.target_networks):
for target_param_name, target_param in target_critic.named_parameters():
source_param = self.ensemble_params[target_param_name][target_ind]
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.target_actor.parameters(), self.network.parameters()
):
target_param.data.copy_(
target_param.data * (1.0 - tau) + source_param.data * tau
)
# ---------- Sampling ----------#
def forward(
self,
cond,
deterministic=False,
reparameterize=False,
):
"""use both pre-trained and online policies"""
q1_ind, q2_ind = self.get_random_indices()
# sample an action from the BC policy
bc_action = super().forward(
cond=cond,
deterministic=True,
network_override=self.bc_policy,
)
# sample an action from the RL policy
rl_action = super().forward(
cond=cond,
deterministic=deterministic,
reparameterize=reparameterize,
)
# compute Q value of BC policy
q_bc_1 = self.critic_networks[q1_ind](cond, bc_action) # (B,)
q_bc_2 = self.critic_networks[q2_ind](cond, bc_action)
q_bc = torch.min(q_bc_1, q_bc_2)
# compute Q value of RL policy
q_rl_1 = self.critic_networks[q1_ind](cond, rl_action)
q_rl_2 = self.critic_networks[q2_ind](cond, rl_action)
q_rl = torch.min(q_rl_1, q_rl_2)
# soft sample or greedy
if deterministic or not self.soft_action_sample:
action = torch.where(
(q_bc > q_rl)[:, None, None],
bc_action,
rl_action,
)
else:
# compute the Q weights with probability proportional to exp(\beta * Q(a))
qw_bc = torch.exp(q_bc * self.soft_action_sample_beta)
qw_rl = torch.exp(q_rl * self.soft_action_sample_beta)
q_weights = torch.softmax(
torch.stack([qw_bc, qw_rl], dim=-1),
dim=-1,
)
# sample according to the weights
q_indices = torch.multinomial(q_weights, 1)
action = torch.where(
(q_indices == 0)[:, None],
bc_action,
rl_action,
)
return action
+131
View File
@@ -0,0 +1,131 @@
"""
Reinforcement learning with prior data (RLPD) for Gaussian policy.
Use ensemble of critics.
"""
import torch
import torch.nn as nn
import logging
from copy import deepcopy
from model.common.gaussian import GaussianModel
log = logging.getLogger(__name__)
class RLPD_Gaussian(GaussianModel):
def __init__(
self,
actor,
critic,
n_critics,
backup_entropy=False,
**kwargs,
):
super().__init__(network=actor, **kwargs)
self.n_critics = n_critics
self.backup_entropy = backup_entropy
# initialize critic networks
self.critic_networks = [
deepcopy(critic).to(self.device) for _ in range(n_critics)
]
self.critic_networks = nn.ModuleList(self.critic_networks)
# initialize target networks
self.target_networks = [
deepcopy(critic).to(self.device) for _ in range(n_critics)
]
self.target_networks = nn.ModuleList(self.target_networks)
# Construct a "stateless" version of one of the models. It is "stateless" in the sense that the parameters are meta Tensors and do not have storage.
base_model = deepcopy(self.critic_networks[0])
self.base_model = base_model.to("meta")
self.ensemble_params, self.ensemble_buffers = torch.func.stack_module_state(
self.critic_networks
)
def critic_wrapper(self, params, buffers, data):
"""for vmap"""
return torch.func.functional_call(self.base_model, (params, buffers), data)
def get_random_indices(self, sz=None, num_ind=2):
"""get num_ind random indices from a set of size sz (used for getting critic targets)"""
if sz is None:
sz = len(self.critic_networks)
perm = torch.randperm(sz)
ind = perm[:num_ind].to(self.device)
return ind
def loss_critic(
self,
obs,
next_obs,
actions,
rewards,
terminated,
gamma,
alpha,
):
# get random critic index
q1_ind, q2_ind = self.get_random_indices()
with torch.no_grad():
next_actions, next_logprobs = self.forward(
cond=next_obs,
deterministic=False,
get_logprob=True,
)
next_q1 = self.target_networks[q1_ind](next_obs, next_actions)
next_q2 = self.target_networks[q2_ind](next_obs, next_actions)
next_q = torch.min(next_q1, next_q2)
# target value
target_q = rewards + gamma * (1 - terminated) * next_q # (B,)
# add entropy term to the target
if self.backup_entropy:
target_q = target_q + gamma * (1 - terminated) * alpha * (
-next_logprobs
)
# run all critics in batch
current_q = torch.vmap(self.critic_wrapper, in_dims=(0, 0, None))(
self.ensemble_params, self.ensemble_buffers, (obs, actions)
) # (n_critics, B)
loss_critic = torch.mean((current_q - target_q[None]) ** 2)
return loss_critic
def loss_actor(self, obs, alpha):
action, logprob = self.forward(
obs,
deterministic=False,
reparameterize=True,
get_logprob=True,
)
current_q = torch.vmap(self.critic_wrapper, in_dims=(0, 0, None))(
self.ensemble_params, self.ensemble_buffers, (obs, action)
) # (n_critics, B)
current_q = current_q.mean(dim=0) + alpha * (-logprob)
loss_actor = -torch.mean(current_q)
return loss_actor
def loss_temperature(self, obs, alpha, target_entropy):
with torch.no_grad():
_, logprob = self.forward(
obs,
deterministic=False,
get_logprob=True,
)
loss_alpha = -torch.mean(alpha * (logprob + target_entropy))
return loss_alpha
def update_target_critic(self, tau):
"""need to use ensemble_params instead of critic_networks"""
for target_ind, target_critic in enumerate(self.target_networks):
for target_param_name, target_param in target_critic.named_parameters():
source_param = self.ensemble_params[target_param_name][target_ind]
target_param.data.copy_(
target_param.data * (1.0 - tau) + source_param.data * tau
)
+88
View File
@@ -0,0 +1,88 @@
"""
Soft Actor Critic (SAC) with Gaussian policy.
"""
import torch
import logging
from copy import deepcopy
import torch.nn.functional as F
from model.common.gaussian import GaussianModel
log = logging.getLogger(__name__)
class SAC_Gaussian(GaussianModel):
def __init__(
self,
actor,
critic,
**kwargs,
):
super().__init__(network=actor, **kwargs)
# initialize doubel critic networks
self.critic = critic.to(self.device)
# initialize double target networks
self.target_critic = deepcopy(self.critic).to(self.device)
def loss_critic(
self,
obs,
next_obs,
actions,
rewards,
terminated,
gamma,
alpha,
):
with torch.no_grad():
next_actions, next_logprobs = self.forward(
cond=next_obs,
deterministic=False,
get_logprob=True,
)
next_q1, next_q2 = self.target_critic(
next_obs,
next_actions,
)
next_q = torch.min(next_q1, next_q2) - alpha * next_logprobs
# target value
target_q = rewards + gamma * next_q * (1 - terminated)
current_q1, current_q2 = self.critic(obs, actions)
loss_critic = F.mse_loss(current_q1, target_q) + F.mse_loss(
current_q2, target_q
)
return loss_critic
def loss_actor(self, obs, alpha):
action, logprob = self.forward(
obs,
deterministic=False,
reparameterize=True,
get_logprob=True,
)
current_q1, current_q2 = self.critic(obs, action)
loss_actor = -torch.min(current_q1, current_q2) + alpha * logprob
return loss_actor.mean()
def loss_temperature(self, obs, alpha, target_entropy):
with torch.no_grad():
_, logprob = self.forward(
obs,
deterministic=False,
get_logprob=True,
)
loss_alpha = -torch.mean(alpha * (logprob + target_entropy))
return loss_alpha
def update_target_critic(self, tau):
for target_param, source_param in zip(
self.target_critic.parameters(), self.critic.parameters()
):
target_param.data.copy_(
target_param.data * (1.0 - tau) + source_param.data * tau
)