squash commits
This commit is contained in:
@@ -1,6 +1,14 @@
|
||||
"""
|
||||
PPO for Gaussian policy.
|
||||
|
||||
To: observation sequence length
|
||||
Ta: action chunk size
|
||||
Do: observation dimension
|
||||
Da: action dimension
|
||||
|
||||
C: image channels
|
||||
H, W: image height and width
|
||||
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
@@ -41,8 +49,10 @@ class PPO_Gaussian(VPG_Gaussian):
|
||||
"""
|
||||
PPO loss
|
||||
|
||||
obs: (B, obs_step, obs_dim)
|
||||
actions: (B, horizon_step, action_dim)
|
||||
obs: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
rgb: (B, To, C, H, W)
|
||||
actions: (B, Ta, Da)
|
||||
returns: (B, )
|
||||
values: (B, )
|
||||
advantages: (B,)
|
||||
|
||||
@@ -29,9 +29,8 @@ class RWR_Gaussian(GaussianModel):
|
||||
|
||||
# override
|
||||
def loss(self, actions, obs, reward_weights):
|
||||
cond = obs
|
||||
B = cond.shape[0]
|
||||
means, scales = self.network(cond)
|
||||
B = len(obs)
|
||||
means, scales = self.network(obs)
|
||||
|
||||
dist = D.Normal(loc=means, scale=scales)
|
||||
log_prob = dist.log_prob(actions.view(B, -1)).mean(-1)
|
||||
@@ -42,16 +41,8 @@ class RWR_Gaussian(GaussianModel):
|
||||
# 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),
|
||||
cond=cond,
|
||||
deterministic=deterministic,
|
||||
randn_clip_value=self.randn_clip_value,
|
||||
)
|
||||
|
||||
+19
-34
@@ -15,26 +15,14 @@ class VPG_Gaussian(GaussianModel):
|
||||
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
|
||||
@@ -44,15 +32,31 @@ class VPG_Gaussian(GaussianModel):
|
||||
for param in self.actor.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
# ---------- Sampling ----------#
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
cond,
|
||||
deterministic=False,
|
||||
use_base_policy=False,
|
||||
):
|
||||
return super().forward(
|
||||
cond=cond,
|
||||
deterministic=deterministic,
|
||||
randn_clip_value=self.randn_clip_value,
|
||||
network_override=self.actor if use_base_policy else None,
|
||||
)
|
||||
|
||||
# ---------- RL training ----------#
|
||||
|
||||
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)
|
||||
B = len(actions)
|
||||
dist = self.forward_train(
|
||||
cond,
|
||||
deterministic=False,
|
||||
@@ -66,22 +70,3 @@ class VPG_Gaussian(GaussianModel):
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
+12
-2
@@ -1,6 +1,14 @@
|
||||
"""
|
||||
PPO for GMM policy.
|
||||
|
||||
To: observation sequence length
|
||||
Ta: action chunk size
|
||||
Do: observation dimension
|
||||
Da: action dimension
|
||||
|
||||
C: image channels
|
||||
H, W: image height and width
|
||||
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
@@ -41,8 +49,10 @@ class PPO_GMM(VPG_GMM):
|
||||
"""
|
||||
PPO loss
|
||||
|
||||
obs: (B, obs_step, obs_dim)
|
||||
actions: (B, horizon_step, action_dim)
|
||||
obs: dict with key state/rgb; more recent obs at the end
|
||||
state: (B, To, Do)
|
||||
rgb: (B, To, C, H, W)
|
||||
actions: (B, Ta, Da)
|
||||
returns: (B, )
|
||||
values: (B, )
|
||||
advantages: (B,)
|
||||
|
||||
+13
-23
@@ -8,36 +8,35 @@ class VPG_GMM(GMMModel):
|
||||
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)
|
||||
|
||||
# ---------- Sampling ----------#
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, cond, deterministic=False):
|
||||
return super().forward(
|
||||
cond=cond,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
|
||||
# ---------- RL training ----------#
|
||||
|
||||
def get_logprobs(
|
||||
self,
|
||||
cond,
|
||||
actions,
|
||||
):
|
||||
B, T, D = actions.shape
|
||||
B = len(actions)
|
||||
dist, entropy, std = self.forward_train(
|
||||
cond.view(B, -1),
|
||||
cond,
|
||||
deterministic=False,
|
||||
)
|
||||
log_prob = dist.log_prob(actions.view(B, -1))
|
||||
@@ -45,12 +44,3 @@ class VPG_GMM(GMMModel):
|
||||
|
||||
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