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:
@@ -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