v0.5 to main (#10)

* v0.5 (#9)

* update idql configs

* update awr configs

* update dipo configs

* update qsm configs

* update dqm configs

* update project version to 0.5.0
This commit is contained in:
Allen Z. Ren
2024-10-07 16:35:13 -04:00
committed by GitHub
parent dd14c5887c
commit e0842e71dc
267 changed files with 6769 additions and 1645 deletions
+187
View File
@@ -0,0 +1,187 @@
"""
Calibrated Conservative Q-Learning (CalQL) for Gaussian policy.
"""
import torch
import torch.nn as nn
import logging
from copy import deepcopy
import numpy as np
import einops
from model.common.gaussian import GaussianModel
log = logging.getLogger(__name__)
class CalQL_Gaussian(GaussianModel):
def __init__(
self,
actor,
critic,
network_path=None,
cql_clip_diff_min=-np.inf,
cql_clip_diff_max=np.inf,
cql_min_q_weight=5.0,
cql_n_actions=10,
**kwargs,
):
super().__init__(network=actor, network_path=None, **kwargs)
self.cql_clip_diff_min = cql_clip_diff_min
self.cql_clip_diff_max = cql_clip_diff_max
self.cql_min_q_weight = cql_min_q_weight
self.cql_n_actions = cql_n_actions
# initialize critic networks
self.critic = critic.to(self.device)
self.target_critic = deepcopy(critic).to(self.device)
# Load pre-trained checkpoint - note we are also loading the pre-trained critic here
if network_path is not None:
checkpoint = torch.load(
network_path,
map_location=self.device,
weights_only=True,
)
self.load_state_dict(
checkpoint["model"],
strict=True,
)
log.info("Loaded actor from %s", network_path)
log.info(
f"Number of network parameters: {sum(p.numel() for p in self.parameters())}"
)
def loss_critic(
self,
obs,
next_obs,
actions,
random_actions,
rewards,
returns,
terminated,
gamma,
alpha,
):
B = len(actions)
# Get initial TD loss
q_data1, q_data2 = self.critic(obs, actions)
with torch.no_grad():
# repeat for action samples
next_obs["state"] = next_obs["state"].repeat_interleave(
self.cql_n_actions, dim=0
)
# Get the next actions and logprobs
next_actions, next_logprobs = self.forward(
next_obs,
deterministic=False,
get_logprob=True,
)
next_q1, next_q2 = self.target_critic(next_obs, next_actions)
next_q = torch.min(next_q1, next_q2)
# Reshape the next_q to match the number of samples
next_q = next_q.view(B, self.cql_n_actions) # (B, n_sample)
next_logprobs = next_logprobs.view(B, self.cql_n_actions) # (B, n_sample)
# Get the max indices over the samples, and index into the next_q and next_log_probs
max_idx = torch.argmax(next_q, dim=1)
next_q = next_q[torch.arange(B), max_idx]
next_logprobs = next_logprobs[torch.arange(B), max_idx]
# Get the target Q values
target_q = rewards + gamma * (1 - terminated) * next_q
# Subtract the entropy bonus
target_q = target_q - alpha * next_logprobs
# TD loss
td_loss_1 = nn.functional.mse_loss(q_data1, target_q)
td_loss_2 = nn.functional.mse_loss(q_data2, target_q)
# Get actions and logprobs
log_rand_pi = 0.5 ** torch.prod(torch.tensor(random_actions.shape[-2:]))
pi_actions, log_pi = self.forward(
obs,
deterministic=False,
reparameterize=False,
get_logprob=True,
) # no gradient
# Random action Q values
n_random_actions = random_actions.shape[1]
obs_sample_state = {
"state": obs["state"].repeat_interleave(n_random_actions, dim=0)
}
random_actions = einops.rearrange(random_actions, "B N H A -> (B N) H A")
# Get the random action Q-values
q_rand_1, q_rand_2 = self.critic(obs_sample_state, random_actions)
q_rand_1 = q_rand_1 - log_rand_pi
q_rand_2 = q_rand_2 - log_rand_pi
# Reshape the random action Q values to match the number of samples
q_rand_1 = q_rand_1.view(B, n_random_actions) # (n_sample, B)
q_rand_2 = q_rand_2.view(B, n_random_actions)
# Policy action Q values
q_pi_1, q_pi_2 = self.critic(obs, pi_actions)
q_pi_1 = q_pi_1 - log_pi
q_pi_2 = q_pi_2 - log_pi
# Ensure calibration w.r.t. value function estimate
q_pi_1 = torch.max(q_pi_1, returns)[:, None] # (B, 1)
q_pi_2 = torch.max(q_pi_2, returns)[:, None] # (B, 1)
cat_q_1 = torch.cat([q_rand_1, q_pi_1], dim=-1) # (B, num_samples+1)
cql_qf1_ood = torch.logsumexp(cat_q_1, dim=-1) # max over num_samples
cat_q_2 = torch.cat([q_rand_2, q_pi_2], dim=-1) # (B, num_samples+1)
cql_qf2_ood = torch.logsumexp(cat_q_2, dim=-1) # sum over num_samples
# Subtract the log likelihood of the data
cql_qf1_diff = torch.clamp(
cql_qf1_ood - q_data1,
min=self.cql_clip_diff_min,
max=self.cql_clip_diff_max,
).mean()
cql_qf2_diff = torch.clamp(
cql_qf2_ood - q_data2,
min=self.cql_clip_diff_min,
max=self.cql_clip_diff_max,
).mean()
cql_min_qf1_loss = cql_qf1_diff * self.cql_min_q_weight
cql_min_qf2_loss = cql_qf2_diff * self.cql_min_q_weight
# Sum the two losses
critic_loss = td_loss_1 + td_loss_2 + cql_min_qf1_loss + cql_min_qf2_loss
return critic_loss
def loss_actor(self, obs, alpha):
action, logprob = self.forward(
obs,
deterministic=False,
reparameterize=True,
get_logprob=True,
)
q1, q2 = self.critic(obs, action)
actor_loss = -torch.min(q1, q2) + alpha * logprob
return actor_loss.mean()
def loss_temperature(self, obs, alpha, target_entropy):
with torch.no_grad():
_, logprob = self.forward(
obs,
deterministic=False,
get_logprob=True,
)
loss_alpha = -torch.mean(alpha * (logprob + target_entropy))
return loss_alpha
def update_target_critic(self, tau):
for target_param, param in zip(
self.target_critic.parameters(), self.critic.parameters()
):
target_param.data.copy_(tau * param.data + (1 - tau) * target_param.data)
+205
View File
@@ -0,0 +1,205 @@
"""
Imitation Bootstrapped Reinforcement Learning (IBRL) for Gaussian policy.
"""
import torch
import torch.nn as nn
import logging
from copy import deepcopy
from model.common.gaussian import GaussianModel
log = logging.getLogger(__name__)
class IBRL_Gaussian(GaussianModel):
def __init__(
self,
actor,
critic,
n_critics,
soft_action_sample=False,
soft_action_sample_beta=0.1,
**kwargs,
):
super().__init__(network=actor, **kwargs)
self.soft_action_sample = soft_action_sample
self.soft_action_sample_beta = soft_action_sample_beta
# Set up target actor
self.target_actor = deepcopy(actor)
# Frozen pre-trained policy
self.bc_policy = deepcopy(actor)
for param in self.bc_policy.parameters():
param.requires_grad = False
# initialize critic networks
self.critic_networks = [
deepcopy(critic).to(self.device) for _ in range(n_critics)
]
self.critic_networks = nn.ModuleList(self.critic_networks)
# initialize target networks
self.target_networks = [
deepcopy(critic).to(self.device) for _ in range(n_critics)
]
self.target_networks = nn.ModuleList(self.target_networks)
# Construct a "stateless" version of one of the models. It is "stateless" in the sense that the parameters are meta Tensors and do not have storage.
base_model = deepcopy(self.critic_networks[0])
self.base_model = base_model.to("meta")
self.ensemble_params, self.ensemble_buffers = torch.func.stack_module_state(
self.critic_networks
)
def critic_wrapper(self, params, buffers, data):
"""for vmap"""
return torch.func.functional_call(self.base_model, (params, buffers), data)
def get_random_indices(self, sz=None, num_ind=2):
"""get num_ind random indices from a set of size sz (used for getting critic targets)"""
if sz is None:
sz = len(self.critic_networks)
perm = torch.randperm(sz)
ind = perm[:num_ind].to(self.device)
return ind
def loss_critic(
self,
obs,
next_obs,
actions,
rewards,
terminated,
gamma,
):
# get random critic index
q1_ind, q2_ind = self.get_random_indices()
with torch.no_grad():
next_actions_bc = super().forward(
cond=next_obs,
deterministic=True,
network_override=self.bc_policy,
)
next_actions_rl = super().forward(
cond=next_obs,
deterministic=False,
network_override=self.target_actor,
)
# get the BC Q value
next_q1_bc = self.target_networks[q1_ind](next_obs, next_actions_bc)
next_q2_bc = self.target_networks[q2_ind](next_obs, next_actions_bc)
next_q_bc = torch.min(next_q1_bc, next_q2_bc)
# get the RL Q value
next_q1_rl = self.target_networks[q1_ind](next_obs, next_actions_rl)
next_q2_rl = self.target_networks[q2_ind](next_obs, next_actions_rl)
next_q_rl = torch.min(next_q1_rl, next_q2_rl)
# take the max Q value
next_q = torch.where(next_q_bc > next_q_rl, next_q_bc, next_q_rl)
# target value
target_q = rewards + gamma * (1 - terminated) * next_q # (B,)
# run all critics in batch
current_q = torch.vmap(self.critic_wrapper, in_dims=(0, 0, None))(
self.ensemble_params, self.ensemble_buffers, (obs, actions)
) # (n_critics, B)
loss_critic = torch.mean((current_q - target_q[None]) ** 2)
return loss_critic
def loss_actor(self, obs):
action = super().forward(
obs,
deterministic=False,
reparameterize=True,
) # use online policy only, also IBRL does not use tanh squashing
current_q = torch.vmap(self.critic_wrapper, in_dims=(0, 0, None))(
self.ensemble_params, self.ensemble_buffers, (obs, action)
) # (n_critics, B)
current_q = current_q.min(
dim=0
).values # unlike RLPD, IBRL uses the min Q value for actor update
loss_actor = -torch.mean(current_q)
return loss_actor
def update_target_critic(self, tau):
"""need to use ensemble_params instead of critic_networks"""
for target_ind, target_critic in enumerate(self.target_networks):
for target_param_name, target_param in target_critic.named_parameters():
source_param = self.ensemble_params[target_param_name][target_ind]
target_param.data.copy_(
target_param.data * (1.0 - tau) + source_param.data * tau
)
def update_target_actor(self, tau):
for target_param, source_param in zip(
self.target_actor.parameters(), self.network.parameters()
):
target_param.data.copy_(
target_param.data * (1.0 - tau) + source_param.data * tau
)
# ---------- Sampling ----------#
def forward(
self,
cond,
deterministic=False,
reparameterize=False,
):
"""use both pre-trained and online policies"""
q1_ind, q2_ind = self.get_random_indices()
# sample an action from the BC policy
bc_action = super().forward(
cond=cond,
deterministic=True,
network_override=self.bc_policy,
)
# sample an action from the RL policy
rl_action = super().forward(
cond=cond,
deterministic=deterministic,
reparameterize=reparameterize,
)
# compute Q value of BC policy
q_bc_1 = self.critic_networks[q1_ind](cond, bc_action) # (B,)
q_bc_2 = self.critic_networks[q2_ind](cond, bc_action)
q_bc = torch.min(q_bc_1, q_bc_2)
# compute Q value of RL policy
q_rl_1 = self.critic_networks[q1_ind](cond, rl_action)
q_rl_2 = self.critic_networks[q2_ind](cond, rl_action)
q_rl = torch.min(q_rl_1, q_rl_2)
# soft sample or greedy
if deterministic or not self.soft_action_sample:
action = torch.where(
(q_bc > q_rl)[:, None, None],
bc_action,
rl_action,
)
else:
# compute the Q weights with probability proportional to exp(\beta * Q(a))
qw_bc = torch.exp(q_bc * self.soft_action_sample_beta)
qw_rl = torch.exp(q_rl * self.soft_action_sample_beta)
q_weights = torch.softmax(
torch.stack([qw_bc, qw_rl], dim=-1),
dim=-1,
)
# sample according to the weights
q_indices = torch.multinomial(q_weights, 1)
action = torch.where(
(q_indices == 0)[:, None],
bc_action,
rl_action,
)
return action
+131
View File
@@ -0,0 +1,131 @@
"""
Reinforcement learning with prior data (RLPD) for Gaussian policy.
Use ensemble of critics.
"""
import torch
import torch.nn as nn
import logging
from copy import deepcopy
from model.common.gaussian import GaussianModel
log = logging.getLogger(__name__)
class RLPD_Gaussian(GaussianModel):
def __init__(
self,
actor,
critic,
n_critics,
backup_entropy=False,
**kwargs,
):
super().__init__(network=actor, **kwargs)
self.n_critics = n_critics
self.backup_entropy = backup_entropy
# initialize critic networks
self.critic_networks = [
deepcopy(critic).to(self.device) for _ in range(n_critics)
]
self.critic_networks = nn.ModuleList(self.critic_networks)
# initialize target networks
self.target_networks = [
deepcopy(critic).to(self.device) for _ in range(n_critics)
]
self.target_networks = nn.ModuleList(self.target_networks)
# Construct a "stateless" version of one of the models. It is "stateless" in the sense that the parameters are meta Tensors and do not have storage.
base_model = deepcopy(self.critic_networks[0])
self.base_model = base_model.to("meta")
self.ensemble_params, self.ensemble_buffers = torch.func.stack_module_state(
self.critic_networks
)
def critic_wrapper(self, params, buffers, data):
"""for vmap"""
return torch.func.functional_call(self.base_model, (params, buffers), data)
def get_random_indices(self, sz=None, num_ind=2):
"""get num_ind random indices from a set of size sz (used for getting critic targets)"""
if sz is None:
sz = len(self.critic_networks)
perm = torch.randperm(sz)
ind = perm[:num_ind].to(self.device)
return ind
def loss_critic(
self,
obs,
next_obs,
actions,
rewards,
terminated,
gamma,
alpha,
):
# get random critic index
q1_ind, q2_ind = self.get_random_indices()
with torch.no_grad():
next_actions, next_logprobs = self.forward(
cond=next_obs,
deterministic=False,
get_logprob=True,
)
next_q1 = self.target_networks[q1_ind](next_obs, next_actions)
next_q2 = self.target_networks[q2_ind](next_obs, next_actions)
next_q = torch.min(next_q1, next_q2)
# target value
target_q = rewards + gamma * (1 - terminated) * next_q # (B,)
# add entropy term to the target
if self.backup_entropy:
target_q = target_q + gamma * (1 - terminated) * alpha * (
-next_logprobs
)
# run all critics in batch
current_q = torch.vmap(self.critic_wrapper, in_dims=(0, 0, None))(
self.ensemble_params, self.ensemble_buffers, (obs, actions)
) # (n_critics, B)
loss_critic = torch.mean((current_q - target_q[None]) ** 2)
return loss_critic
def loss_actor(self, obs, alpha):
action, logprob = self.forward(
obs,
deterministic=False,
reparameterize=True,
get_logprob=True,
)
current_q = torch.vmap(self.critic_wrapper, in_dims=(0, 0, None))(
self.ensemble_params, self.ensemble_buffers, (obs, action)
) # (n_critics, B)
current_q = current_q.mean(dim=0) + alpha * (-logprob)
loss_actor = -torch.mean(current_q)
return loss_actor
def loss_temperature(self, obs, alpha, target_entropy):
with torch.no_grad():
_, logprob = self.forward(
obs,
deterministic=False,
get_logprob=True,
)
loss_alpha = -torch.mean(alpha * (logprob + target_entropy))
return loss_alpha
def update_target_critic(self, tau):
"""need to use ensemble_params instead of critic_networks"""
for target_ind, target_critic in enumerate(self.target_networks):
for target_param_name, target_param in target_critic.named_parameters():
source_param = self.ensemble_params[target_param_name][target_ind]
target_param.data.copy_(
target_param.data * (1.0 - tau) + source_param.data * tau
)
+88
View File
@@ -0,0 +1,88 @@
"""
Soft Actor Critic (SAC) with Gaussian policy.
"""
import torch
import logging
from copy import deepcopy
import torch.nn.functional as F
from model.common.gaussian import GaussianModel
log = logging.getLogger(__name__)
class SAC_Gaussian(GaussianModel):
def __init__(
self,
actor,
critic,
**kwargs,
):
super().__init__(network=actor, **kwargs)
# initialize doubel critic networks
self.critic = critic.to(self.device)
# initialize double target networks
self.target_critic = deepcopy(self.critic).to(self.device)
def loss_critic(
self,
obs,
next_obs,
actions,
rewards,
terminated,
gamma,
alpha,
):
with torch.no_grad():
next_actions, next_logprobs = self.forward(
cond=next_obs,
deterministic=False,
get_logprob=True,
)
next_q1, next_q2 = self.target_critic(
next_obs,
next_actions,
)
next_q = torch.min(next_q1, next_q2) - alpha * next_logprobs
# target value
target_q = rewards + gamma * next_q * (1 - terminated)
current_q1, current_q2 = self.critic(obs, actions)
loss_critic = F.mse_loss(current_q1, target_q) + F.mse_loss(
current_q2, target_q
)
return loss_critic
def loss_actor(self, obs, alpha):
action, logprob = self.forward(
obs,
deterministic=False,
reparameterize=True,
get_logprob=True,
)
current_q1, current_q2 = self.critic(obs, action)
loss_actor = -torch.min(current_q1, current_q2) + alpha * logprob
return loss_actor.mean()
def loss_temperature(self, obs, alpha, target_entropy):
with torch.no_grad():
_, logprob = self.forward(
obs,
deterministic=False,
get_logprob=True,
)
loss_alpha = -torch.mean(alpha * (logprob + target_entropy))
return loss_alpha
def update_target_critic(self, tau):
for target_param, source_param in zip(
self.target_critic.parameters(), self.critic.parameters()
):
target_param.data.copy_(
target_param.data * (1.0 - tau) + source_param.data * tau
)