support varying img size

This commit is contained in:
allenzren
2024-09-16 17:55:31 -04:00
parent 64595baca9
commit 1aaa6c2302
18 changed files with 131 additions and 81 deletions
+6 -5
View File
@@ -244,15 +244,16 @@ class TrainPPODiffusionAgent(TrainPPOAgent):
.float()
.to(self.device)
}
with torch.no_grad():
next_value = (
self.model.critic(obs_venv_ts).reshape(1, -1).cpu().numpy()
)
advantages_trajs = np.zeros_like(reward_trajs)
lastgaelam = 0
for t in reversed(range(self.n_steps)):
if t == self.n_steps - 1:
nextvalues = next_value
nextvalues = (
self.model.critic(obs_venv_ts)
.reshape(1, -1)
.cpu()
.numpy()
)
else:
nextvalues = values_trajs[t + 1]
nonterminal = 1.0 - dones_trajs[t]
@@ -240,18 +240,16 @@ class TrainPPOImgDiffusionAgent(TrainPPODiffusionAgent):
key: torch.from_numpy(obs_venv[key]).float().to(self.device)
for key in self.obs_dims
}
with torch.no_grad():
next_value = (
self.model.critic(obs_venv_ts, no_augment=True)
.reshape(1, -1)
.cpu()
.numpy()
)
advantages_trajs = np.zeros_like(reward_trajs)
lastgaelam = 0
for t in reversed(range(self.n_steps)):
if t == self.n_steps - 1:
nextvalues = next_value
nextvalues = (
self.model.critic(obs_venv_ts, no_augment=True)
.reshape(1, -1)
.cpu()
.numpy()
)
else:
nextvalues = values_trajs[t + 1]
nonterminal = 1.0 - dones_trajs[t]
@@ -220,15 +220,16 @@ class TrainPPOExactDiffusionAgent(TrainPPODiffusionAgent):
.float()
.to(self.device)
}
with torch.no_grad():
next_value = (
self.model.critic(obs_venv_ts).reshape(1, -1).cpu().numpy()
)
advantages_trajs = np.zeros_like(reward_trajs)
lastgaelam = 0
for t in reversed(range(self.n_steps)):
if t == self.n_steps - 1:
nextvalues = next_value
nextvalues = (
self.model.critic(obs_venv_ts)
.reshape(1, -1)
.cpu()
.numpy()
)
else:
nextvalues = values_trajs[t + 1]
nonterminal = 1.0 - dones_trajs[t]
@@ -241,10 +242,7 @@ class TrainPPOExactDiffusionAgent(TrainPPODiffusionAgent):
# A = delta_t + gamma*lamdba*delta_{t+1} + ...
advantages_trajs[t] = lastgaelam = (
delta
+ self.gamma
* self.gae_lambda
* nonterminal
* lastgaelam
+ self.gamma * self.gae_lambda * nonterminal * lastgaelam
)
returns_trajs = advantages_trajs + values_trajs
+6 -5
View File
@@ -209,15 +209,16 @@ class TrainPPOGaussianAgent(TrainPPOAgent):
.float()
.to(self.device)
}
with torch.no_grad():
next_value = (
self.model.critic(obs_venv_ts).reshape(1, -1).cpu().numpy()
)
advantages_trajs = np.zeros_like(reward_trajs)
lastgaelam = 0
for t in reversed(range(self.n_steps)):
if t == self.n_steps - 1:
nextvalues = next_value
nextvalues = (
self.model.critic(obs_venv_ts)
.reshape(1, -1)
.cpu()
.numpy()
)
else:
nextvalues = values_trajs[t + 1]
nonterminal = 1.0 - dones_trajs[t]
@@ -228,18 +228,16 @@ class TrainPPOImgGaussianAgent(TrainPPOGaussianAgent):
key: torch.from_numpy(obs_venv[key]).float().to(self.device)
for key in self.obs_dims
}
with torch.no_grad():
next_value = (
self.model.critic(obs_venv_ts, no_augment=True)
.reshape(1, -1)
.cpu()
.numpy()
)
advantages_trajs = np.zeros_like(reward_trajs)
lastgaelam = 0
for t in reversed(range(self.n_steps)):
if t == self.n_steps - 1:
nextvalues = next_value
nextvalues = (
self.model.critic(obs_venv_ts, no_augment=True)
.reshape(1, -1)
.cpu()
.numpy()
)
else:
nextvalues = values_trajs[t + 1]
nonterminal = 1.0 - dones_trajs[t]