add minor docs to diffusion classes and clean up some args

This commit is contained in:
allenzren
2024-09-17 16:26:25 -04:00
parent ef5b14f820
commit bc52beca1e
8 changed files with 33 additions and 98 deletions
-5
View File
@@ -16,7 +16,6 @@ class RWR_Gaussian(GaussianModel):
def __init__(
self,
actor,
randn_clip_value=10,
**kwargs,
):
super().__init__(network=actor, **kwargs)
@@ -24,9 +23,6 @@ class RWR_Gaussian(GaussianModel):
# 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):
B = len(obs)
@@ -44,6 +40,5 @@ class RWR_Gaussian(GaussianModel):
actions = super().forward(
cond=cond,
deterministic=deterministic,
randn_clip_value=self.randn_clip_value,
)
return actions
-3
View File
@@ -15,11 +15,9 @@ class VPG_Gaussian(GaussianModel):
self,
actor,
critic,
randn_clip_value=10,
**kwargs,
):
super().__init__(network=actor, **kwargs)
self.randn_clip_value = randn_clip_value
# Value function for obs - simple MLP
self.critic = critic.to(self.device)
@@ -44,7 +42,6 @@ class VPG_Gaussian(GaussianModel):
return super().forward(
cond=cond,
deterministic=deterministic,
randn_clip_value=self.randn_clip_value,
network_override=self.actor if use_base_policy else None,
)