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)
|
||||
|
||||
Reference in New Issue
Block a user