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
+6 -3
View File
@@ -18,6 +18,7 @@ class GaussianModel(torch.nn.Module):
horizon_steps,
network_path=None,
device="cuda:0",
randn_clip_value=10,
):
super().__init__()
self.device = device
@@ -36,6 +37,9 @@ class GaussianModel(torch.nn.Module):
)
self.horizon_steps = horizon_steps
# 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
def loss(
self,
true_action,
@@ -75,7 +79,6 @@ class GaussianModel(torch.nn.Module):
self,
cond,
deterministic=False,
randn_clip_value=10,
network_override=None,
):
B = len(cond["state"]) if "state" in cond else len(cond["rgb"])
@@ -87,7 +90,7 @@ class GaussianModel(torch.nn.Module):
)
sampled_action = dist.sample()
sampled_action.clamp_(
dist.loc - randn_clip_value * dist.scale,
dist.loc + randn_clip_value * dist.scale,
dist.loc - self.randn_clip_value * dist.scale,
dist.loc + self.randn_clip_value * dist.scale,
)
return sampled_action.view(B, T, -1)