Remove unused sampling_weights.

This commit is contained in:
Eloi Alonso
2023-09-18 16:12:04 +02:00
parent 42bbabee7c
commit ac6be401fe
3 changed files with 9 additions and 23 deletions
+4 -15
View File
@@ -61,22 +61,11 @@ class EpisodesDataset:
self.newly_modified_episodes.add(episode_id)
return episode_id
def sample_batch(self, batch_num_samples: int, sequence_length: int, weights: Optional[Tuple[float]] = None, sample_from_start: bool = True) -> Batch:
return self._collate_episodes_segments(self._sample_episodes_segments(batch_num_samples, sequence_length, weights, sample_from_start))
def _sample_episodes_segments(self, batch_num_samples: int, sequence_length: int, weights: Optional[Tuple[float]], sample_from_start: bool) -> List[Episode]:
num_episodes = len(self.episodes)
num_weights = len(weights) if weights is not None else 0
if num_weights < num_episodes:
weights = [1] * num_episodes
else:
assert all([0 <= x <= 1 for x in weights]) and sum(weights) == 1
sizes = [num_episodes // num_weights + (num_episodes % num_weights) * (i == num_weights - 1) for i in range(num_weights)]
weights = [w / s for (w, s) in zip(weights, sizes) for _ in range(s)]
sampled_episodes = random.choices(self.episodes, k=batch_num_samples, weights=weights)
def sample_batch(self, batch_num_samples: int, sequence_length: int, sample_from_start: bool = True) -> Batch:
return self._collate_episodes_segments(self._sample_episodes_segments(batch_num_samples, sequence_length, sample_from_start))
def _sample_episodes_segments(self, batch_num_samples: int, sequence_length: int, sample_from_start: bool) -> List[Episode]:
sampled_episodes = random.choices(self.episodes, k=batch_num_samples)
sampled_episodes_segments = []
for sampled_episode in sampled_episodes:
if sample_from_start:
+5 -7
View File
@@ -134,30 +134,28 @@ class Trainer:
cfg_world_model = self.cfg.training.world_model
cfg_actor_critic = self.cfg.training.actor_critic
w = self.cfg.training.sampling_weights
if epoch > cfg_tokenizer.start_after_epochs:
metrics_tokenizer = self.train_component(self.agent.tokenizer, self.optimizer_tokenizer, sequence_length=1, sample_from_start=True, sampling_weights=w, **cfg_tokenizer)
metrics_tokenizer = self.train_component(self.agent.tokenizer, self.optimizer_tokenizer, sequence_length=1, sample_from_start=True, **cfg_tokenizer)
self.agent.tokenizer.eval()
if epoch > cfg_world_model.start_after_epochs:
metrics_world_model = self.train_component(self.agent.world_model, self.optimizer_world_model, sequence_length=self.cfg.common.sequence_length, sample_from_start=True, sampling_weights=w, tokenizer=self.agent.tokenizer, **cfg_world_model)
metrics_world_model = self.train_component(self.agent.world_model, self.optimizer_world_model, sequence_length=self.cfg.common.sequence_length, sample_from_start=True, tokenizer=self.agent.tokenizer, **cfg_world_model)
self.agent.world_model.eval()
if epoch > cfg_actor_critic.start_after_epochs:
metrics_actor_critic = self.train_component(self.agent.actor_critic, self.optimizer_actor_critic, sequence_length=1 + self.cfg.training.actor_critic.burn_in, sample_from_start=False, sampling_weights=w, tokenizer=self.agent.tokenizer, world_model=self.agent.world_model, **cfg_actor_critic)
metrics_actor_critic = self.train_component(self.agent.actor_critic, self.optimizer_actor_critic, sequence_length=1 + self.cfg.training.actor_critic.burn_in, sample_from_start=False, tokenizer=self.agent.tokenizer, world_model=self.agent.world_model, **cfg_actor_critic)
self.agent.actor_critic.eval()
return [{'epoch': epoch, **metrics_tokenizer, **metrics_world_model, **metrics_actor_critic}]
def train_component(self, component: nn.Module, optimizer: torch.optim.Optimizer, steps_per_epoch: int, batch_num_samples: int, grad_acc_steps: int, max_grad_norm: Optional[float], sequence_length: int, sampling_weights: Optional[Tuple[float]], sample_from_start: bool, **kwargs_loss: Any) -> Dict[str, float]:
def train_component(self, component: nn.Module, optimizer: torch.optim.Optimizer, steps_per_epoch: int, batch_num_samples: int, grad_acc_steps: int, max_grad_norm: Optional[float], sequence_length: int, sample_from_start: bool, **kwargs_loss: Any) -> Dict[str, float]:
loss_total_epoch = 0.0
intermediate_losses = defaultdict(float)
for _ in tqdm(range(steps_per_epoch), desc=f"Training {str(component)}", file=sys.stdout):
optimizer.zero_grad()
for _ in range(grad_acc_steps):
batch = self.train_dataset.sample_batch(batch_num_samples, sequence_length, sampling_weights, sample_from_start)
batch = self.train_dataset.sample_batch(batch_num_samples, sequence_length, sample_from_start)
batch = self._to_device(batch)
losses = component.compute_loss(batch, **kwargs_loss) / grad_acc_steps