add minor docs to diffusion classes and clean up some args
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user