fix link and generate
This commit is contained in:
@@ -233,9 +233,12 @@ class MoECTS:
|
||||
# RND loss
|
||||
mean_rnd_loss = 0 if self.rnd else None
|
||||
|
||||
# Get mini batch generator
|
||||
generator = self.storage.mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
|
||||
data = list(generator)
|
||||
# Reuse the exact same shuffled samples for both update phases without
|
||||
# materializing every epoch and mini-batch on the GPU at once.
|
||||
batch_indices = self.storage.generate_mini_batch_indices()
|
||||
generator = self.storage.mini_batch_generator(
|
||||
self.num_mini_batches, self.num_learning_epochs, batch_indices=batch_indices
|
||||
)
|
||||
|
||||
# Iterate over batches
|
||||
teacher_samples = self.teacher_num_envs * self.storage.num_transitions_per_env // self.num_mini_batches
|
||||
@@ -251,7 +254,7 @@ class MoECTS:
|
||||
old_sigma_batch,
|
||||
hidden_states_batch,
|
||||
masks_batch,
|
||||
) in data:
|
||||
) in generator:
|
||||
original_batch_size = obs_batch.batch_size[0]
|
||||
|
||||
# Check if we should normalize advantages per mini batch
|
||||
@@ -377,6 +380,9 @@ class MoECTS:
|
||||
if mean_rnd_loss is not None:
|
||||
mean_rnd_loss += rnd_loss.item()
|
||||
|
||||
generator = self.storage.mini_batch_generator(
|
||||
self.num_mini_batches, self.num_learning_epochs, batch_indices=batch_indices
|
||||
)
|
||||
for (
|
||||
obs_batch,
|
||||
actions_batch,
|
||||
@@ -388,7 +394,7 @@ class MoECTS:
|
||||
old_sigma_batch,
|
||||
hidden_states_batch,
|
||||
masks_batch,
|
||||
) in data:
|
||||
) in generator:
|
||||
# Student encoder loss
|
||||
obs_a_batch = self.policy.get_actor_obs(obs_batch)
|
||||
obs_a_batch = self.policy.actor_obs_normalizer(obs_a_batch)
|
||||
|
||||
@@ -127,7 +127,25 @@ class RolloutStorageCTS:
|
||||
yield self.observations[i], self.actions[i], self.privileged_actions[i], self.dones[i]
|
||||
|
||||
# For reinforcement learning with feedforward networks
|
||||
def mini_batch_generator(self, num_mini_batches: int, num_epochs: int = 8) -> Generator:
|
||||
def generate_mini_batch_indices(self) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Generate the teacher and student permutations shared by both update phases."""
|
||||
if self.training_type != "rl":
|
||||
raise ValueError("This function is only available for reinforcement learning training.")
|
||||
|
||||
teacher_samples_num = self.teacher_num_envs * self.num_transitions_per_env
|
||||
student_samples_num = self.student_num_envs * self.num_transitions_per_env
|
||||
teacher_indices = torch.randperm(teacher_samples_num, requires_grad=False, device=self.device)
|
||||
student_indices = teacher_samples_num + torch.randperm(
|
||||
student_samples_num, requires_grad=False, device=self.device
|
||||
)
|
||||
return teacher_indices, student_indices
|
||||
|
||||
def mini_batch_generator(
|
||||
self,
|
||||
num_mini_batches: int,
|
||||
num_epochs: int = 8,
|
||||
batch_indices: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> Generator:
|
||||
if self.training_type != "rl":
|
||||
raise ValueError("This function is only available for reinforcement learning training.")
|
||||
|
||||
@@ -136,8 +154,10 @@ class RolloutStorageCTS:
|
||||
student_samples_num = self.student_num_envs * self.num_transitions_per_env
|
||||
teacher_mini_batch_size = teacher_samples_num // num_mini_batches
|
||||
student_mini_batch_size = student_samples_num // num_mini_batches
|
||||
teacher_indices = torch.randperm(teacher_samples_num, requires_grad=False, device=self.device)
|
||||
student_indices = teacher_samples_num + torch.randperm(student_samples_num, requires_grad=False, device=self.device)
|
||||
if batch_indices is None:
|
||||
teacher_indices, student_indices = self.generate_mini_batch_indices()
|
||||
else:
|
||||
teacher_indices, student_indices = batch_indices
|
||||
|
||||
# Core
|
||||
observations = self.observations.transpose(0, 1).flatten(0, 1)
|
||||
@@ -204,4 +224,4 @@ class RolloutStorageCTS:
|
||||
return NotImplementedError("CTS rollout storage does not support RNNs yet.")
|
||||
|
||||
def _save_hidden_states(self, hidden_states: tuple[HiddenState, HiddenState]) -> None:
|
||||
return NotImplementedError("CTS rollout storage does not support RNNs yet.")
|
||||
return NotImplementedError("CTS rollout storage does not support RNNs yet.")
|
||||
|
||||
Reference in New Issue
Block a user