squash commits

This commit is contained in:
allenzren
2024-09-11 21:09:17 -04:00
parent 8ce0aa1485
commit 2ddf63b8f5
200 changed files with 1240 additions and 1186 deletions
+12 -2
View File
@@ -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,)
+3 -12
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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,
)