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:
+26
-26
@@ -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):
|
||||
|
||||
@@ -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
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
)
|
||||
@@ -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
|
||||
)
|
||||
Reference in New Issue
Block a user