release
This commit is contained in:
@@ -0,0 +1,30 @@
|
||||
"""
|
||||
Advantage-weighted regression (AWR) for Gaussian policy.
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
import logging
|
||||
from model.rl.gaussian_rwr import RWR_Gaussian
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AWR_Gaussian(RWR_Gaussian):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
actor,
|
||||
critic,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(actor=actor, **kwargs)
|
||||
self.critic = critic.to(self.device)
|
||||
|
||||
def loss_critic(self, obs, advantages):
|
||||
# get advantage
|
||||
adv = self.critic(obs)
|
||||
|
||||
# Update critic
|
||||
loss_critic = torch.mean((adv - advantages) ** 2)
|
||||
return loss_critic
|
||||
@@ -0,0 +1,120 @@
|
||||
"""
|
||||
PPO for Gaussian policy.
|
||||
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
import torch
|
||||
from model.rl.gaussian_vpg import VPG_Gaussian
|
||||
|
||||
|
||||
class PPO_Gaussian(VPG_Gaussian):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
clip_ploss_coef: float,
|
||||
clip_vloss_coef: Optional[float] = None,
|
||||
norm_adv: Optional[bool] = True,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
# Whether to normalize advantages within batch
|
||||
self.norm_adv = norm_adv
|
||||
|
||||
# Clipping value for policy loss
|
||||
self.clip_ploss_coef = clip_ploss_coef
|
||||
|
||||
# Clipping value for value loss
|
||||
self.clip_vloss_coef = clip_vloss_coef
|
||||
|
||||
def loss(
|
||||
self,
|
||||
obs,
|
||||
actions,
|
||||
returns,
|
||||
oldvalues,
|
||||
advantages,
|
||||
oldlogprobs,
|
||||
use_bc_loss=False,
|
||||
):
|
||||
"""
|
||||
PPO loss
|
||||
|
||||
obs: (B, obs_step, obs_dim)
|
||||
actions: (B, horizon_step, action_dim)
|
||||
returns: (B, )
|
||||
values: (B, )
|
||||
advantages: (B,)
|
||||
oldlogprobs: (B, )
|
||||
"""
|
||||
newlogprobs, entropy, std = self.get_logprobs(obs, actions)
|
||||
newlogprobs = newlogprobs.clamp(min=-5, max=2)
|
||||
oldlogprobs = oldlogprobs.clamp(min=-5, max=2)
|
||||
entropy_loss = -entropy
|
||||
|
||||
# get ratio
|
||||
logratio = newlogprobs - oldlogprobs
|
||||
ratio = logratio.exp()
|
||||
|
||||
# get kl difference and whether value clipped
|
||||
with torch.no_grad():
|
||||
approx_kl = ((ratio - 1) - logratio).nanmean()
|
||||
clipfrac = (
|
||||
((ratio - 1.0).abs() > self.clip_ploss_coef).float().mean().item()
|
||||
)
|
||||
|
||||
# normalize advantages
|
||||
if self.norm_adv:
|
||||
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
|
||||
|
||||
# Policy loss with clipping
|
||||
pg_loss1 = -advantages * ratio
|
||||
pg_loss2 = -advantages * torch.clamp(
|
||||
ratio, 1 - self.clip_ploss_coef, 1 + self.clip_ploss_coef
|
||||
)
|
||||
pg_loss = torch.max(pg_loss1, pg_loss2).mean()
|
||||
|
||||
# Value loss optionally with clipping
|
||||
newvalues = self.critic(obs).view(-1)
|
||||
if self.clip_vloss_coef is not None:
|
||||
v_loss_unclipped = (newvalues - returns) ** 2
|
||||
v_clipped = oldvalues + torch.clamp(
|
||||
newvalues - oldvalues,
|
||||
-self.clip_vloss_coef,
|
||||
self.clip_vloss_coef,
|
||||
)
|
||||
v_loss_clipped = (v_clipped - returns) ** 2
|
||||
v_loss_max = torch.max(v_loss_unclipped, v_loss_clipped)
|
||||
v_loss = 0.5 * v_loss_max.mean()
|
||||
else:
|
||||
v_loss = 0.5 * ((newvalues - returns) ** 2).mean()
|
||||
|
||||
bc_loss = 0.0
|
||||
if use_bc_loss:
|
||||
# See Eqn. 2 of https://arxiv.org/pdf/2403.03949.pdf
|
||||
# Give a reward for maximizing probability of teacher policy's action with current policy.
|
||||
# Actions are chosen along trajectory induced by current policy.
|
||||
|
||||
# Get counterfactual teacher actions
|
||||
samples = self.forward(
|
||||
cond=obs.float()
|
||||
.unsqueeze(1)
|
||||
.to(self.device), # B x horizon=1 x obs_dim
|
||||
deterministic=False,
|
||||
use_base_policy=True,
|
||||
)
|
||||
# Get logprobs of teacher actions under this policy
|
||||
bc_logprobs, _, _ = self.get_logprobs(obs, samples, use_base_policy=False)
|
||||
bc_logprobs = bc_logprobs.clamp(min=-5, max=2)
|
||||
bc_loss = -bc_logprobs.mean()
|
||||
return (
|
||||
pg_loss,
|
||||
entropy_loss,
|
||||
v_loss,
|
||||
clipfrac,
|
||||
approx_kl.item(),
|
||||
ratio.mean().item(),
|
||||
bc_loss,
|
||||
std.item(),
|
||||
)
|
||||
@@ -0,0 +1,58 @@
|
||||
"""
|
||||
Reward-weighted regression (RWR) for Gaussian policy.
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
import logging
|
||||
from model.common.gaussian import GaussianModel
|
||||
import torch.distributions as D
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class RWR_Gaussian(GaussianModel):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
actor,
|
||||
randn_clip_value=10,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(network=actor, **kwargs)
|
||||
|
||||
# assign actor
|
||||
self.actor = self.network
|
||||
|
||||
# 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
|
||||
|
||||
# override
|
||||
def loss(self, actions, obs, reward_weights):
|
||||
cond = obs
|
||||
B = cond.shape[0]
|
||||
means, scales = self.network(cond)
|
||||
|
||||
dist = D.Normal(loc=means, scale=scales)
|
||||
log_prob = dist.log_prob(actions.view(B, -1)).mean(-1)
|
||||
log_prob = log_prob * reward_weights
|
||||
log_prob = -log_prob.mean()
|
||||
return log_prob
|
||||
|
||||
# override
|
||||
@torch.no_grad()
|
||||
def forward(self, cond, deterministic=False, **kwargs):
|
||||
"""
|
||||
Args:
|
||||
cond: (batch_size, horizon, obs_dim)
|
||||
|
||||
Return:
|
||||
actions: (batch_size, horizon_steps, transition_dim)
|
||||
"""
|
||||
B = cond.shape[0]
|
||||
actions = super().forward(
|
||||
cond=cond.view(B, -1),
|
||||
deterministic=deterministic,
|
||||
randn_clip_value=self.randn_clip_value,
|
||||
)
|
||||
return actions
|
||||
@@ -0,0 +1,87 @@
|
||||
"""
|
||||
Policy gradient for Gaussian policy
|
||||
|
||||
"""
|
||||
|
||||
import torch
|
||||
from copy import deepcopy
|
||||
import logging
|
||||
from model.common.gaussian import GaussianModel
|
||||
|
||||
|
||||
class VPG_Gaussian(GaussianModel):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
actor,
|
||||
critic,
|
||||
cond_steps=1,
|
||||
randn_clip_value=10,
|
||||
network_path=None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(network=actor, **kwargs)
|
||||
self.cond_steps = cond_steps
|
||||
self.randn_clip_value = randn_clip_value
|
||||
|
||||
# Value function for obs - simple MLP
|
||||
self.critic = critic.to(self.device)
|
||||
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=False,
|
||||
)
|
||||
logging.info("Loaded actor from %s", network_path)
|
||||
|
||||
# Re-name network to actor
|
||||
self.actor_ft = actor
|
||||
|
||||
# Save a copy of original actor
|
||||
self.actor = deepcopy(actor)
|
||||
for param in self.actor.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def get_logprobs(
|
||||
self,
|
||||
cond,
|
||||
actions,
|
||||
use_base_policy=False,
|
||||
):
|
||||
B, T, D = actions.shape
|
||||
if not isinstance(cond, dict):
|
||||
cond = cond.view(B, -1)
|
||||
dist = self.forward_train(
|
||||
cond,
|
||||
deterministic=False,
|
||||
network_override=self.actor if use_base_policy else None,
|
||||
)
|
||||
log_prob = dist.log_prob(actions.view(B, -1))
|
||||
log_prob = log_prob.mean(-1)
|
||||
entropy = dist.entropy().mean()
|
||||
std = dist.scale.mean()
|
||||
return log_prob, entropy, std
|
||||
|
||||
def loss(self, obs, actions, reward):
|
||||
raise NotImplementedError
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
cond,
|
||||
deterministic=False,
|
||||
use_base_policy=False,
|
||||
):
|
||||
if isinstance(cond, dict):
|
||||
B = cond["state"].shape[0]
|
||||
else:
|
||||
B = cond.shape[0]
|
||||
cond = cond.view(B, -1)
|
||||
return super().forward(
|
||||
cond=cond,
|
||||
deterministic=deterministic,
|
||||
randn_clip_value=self.randn_clip_value,
|
||||
network_override=self.actor if use_base_policy else None,
|
||||
)
|
||||
@@ -0,0 +1,102 @@
|
||||
"""
|
||||
PPO for GMM policy.
|
||||
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
import torch
|
||||
from model.rl.gmm_vpg import VPG_GMM
|
||||
|
||||
|
||||
class PPO_GMM(VPG_GMM):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
clip_ploss_coef: float,
|
||||
clip_vloss_coef: Optional[float] = None,
|
||||
norm_adv: Optional[bool] = True,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
# Whether to normalize advantages within batch
|
||||
self.norm_adv = norm_adv
|
||||
|
||||
# Clipping value for policy loss
|
||||
self.clip_ploss_coef = clip_ploss_coef
|
||||
|
||||
# Clipping value for value loss
|
||||
self.clip_vloss_coef = clip_vloss_coef
|
||||
|
||||
def loss(
|
||||
self,
|
||||
obs,
|
||||
actions,
|
||||
returns,
|
||||
oldvalues,
|
||||
advantages,
|
||||
oldlogprobs,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
PPO loss
|
||||
|
||||
obs: (B, obs_step, obs_dim)
|
||||
actions: (B, horizon_step, action_dim)
|
||||
returns: (B, )
|
||||
values: (B, )
|
||||
advantages: (B,)
|
||||
oldlogprobs: (B, )
|
||||
"""
|
||||
newlogprobs, entropy, std = self.get_logprobs(obs, actions)
|
||||
newlogprobs = newlogprobs.clamp(min=-5, max=2)
|
||||
oldlogprobs = oldlogprobs.clamp(min=-5, max=2)
|
||||
entropy_loss = -entropy.mean()
|
||||
|
||||
# get ratio
|
||||
logratio = newlogprobs - oldlogprobs
|
||||
ratio = logratio.exp()
|
||||
|
||||
# get kl difference and whether value clipped
|
||||
with torch.no_grad():
|
||||
approx_kl = ((ratio - 1) - logratio).nanmean()
|
||||
clipfrac = (
|
||||
((ratio - 1.0).abs() > self.clip_ploss_coef).float().mean().item()
|
||||
)
|
||||
|
||||
# normalize advantages
|
||||
if self.norm_adv:
|
||||
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
|
||||
|
||||
# Policy loss with clipping
|
||||
pg_loss1 = -advantages * ratio
|
||||
pg_loss2 = -advantages * torch.clamp(
|
||||
ratio, 1 - self.clip_ploss_coef, 1 + self.clip_ploss_coef
|
||||
)
|
||||
pg_loss = torch.max(pg_loss1, pg_loss2).mean()
|
||||
|
||||
# Value loss optionally with clipping
|
||||
newvalues = self.critic(obs).view(-1)
|
||||
if self.clip_vloss_coef is not None:
|
||||
v_loss_unclipped = (newvalues - returns) ** 2
|
||||
v_clipped = oldvalues + torch.clamp(
|
||||
newvalues - oldvalues,
|
||||
-self.clip_vloss_coef,
|
||||
self.clip_vloss_coef,
|
||||
)
|
||||
v_loss_clipped = (v_clipped - returns) ** 2
|
||||
v_loss_max = torch.max(v_loss_unclipped, v_loss_clipped)
|
||||
v_loss = 0.5 * v_loss_max.mean()
|
||||
else:
|
||||
v_loss = 0.5 * ((newvalues - returns) ** 2).mean()
|
||||
bc_loss = 0
|
||||
return (
|
||||
pg_loss,
|
||||
entropy_loss,
|
||||
v_loss,
|
||||
clipfrac,
|
||||
approx_kl.item(),
|
||||
ratio.mean().item(),
|
||||
bc_loss,
|
||||
std.item(),
|
||||
)
|
||||
@@ -0,0 +1,56 @@
|
||||
import torch
|
||||
import logging
|
||||
from model.common.gmm import GMMModel
|
||||
|
||||
|
||||
class VPG_GMM(GMMModel):
|
||||
def __init__(
|
||||
self,
|
||||
actor,
|
||||
critic,
|
||||
cond_steps=1,
|
||||
network_path=None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(network=actor, **kwargs)
|
||||
self.cond_steps = cond_steps
|
||||
|
||||
# Re-name network to actor
|
||||
self.actor_ft = actor
|
||||
|
||||
# Value function for obs - simple MLP
|
||||
self.critic = critic.to(self.device)
|
||||
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=False,
|
||||
)
|
||||
logging.info("Loaded actor from %s", network_path)
|
||||
|
||||
def get_logprobs(
|
||||
self,
|
||||
cond,
|
||||
actions,
|
||||
):
|
||||
B, T, D = actions.shape
|
||||
dist, entropy, std = self.forward_train(
|
||||
cond.view(B, -1),
|
||||
deterministic=False,
|
||||
)
|
||||
log_prob = dist.log_prob(actions.view(B, -1))
|
||||
return log_prob, entropy, std
|
||||
|
||||
def loss(self, obs, chains, reward):
|
||||
raise NotImplementedError
|
||||
|
||||
# override to diffuse over action only
|
||||
@torch.no_grad()
|
||||
def forward(self, cond, deterministic=False):
|
||||
B = cond.shape[0]
|
||||
return super().forward(
|
||||
cond=cond.view(B, -1),
|
||||
deterministic=deterministic,
|
||||
)
|
||||
Reference in New Issue
Block a user