From ac6be401fed2b6176c9ce0cf1dc10e376c9d740d Mon Sep 17 00:00:00 2001 From: Eloi Alonso Date: Mon, 18 Sep 2023 16:11:06 +0200 Subject: [PATCH] Remove unused sampling_weights. --- config/trainer.yaml | 1 - src/dataset.py | 19 ++++--------------- src/trainer.py | 12 +++++------- 3 files changed, 9 insertions(+), 23 deletions(-) diff --git a/config/trainer.yaml b/config/trainer.yaml index 3890af1..9e14d17 100644 --- a/config/trainer.yaml +++ b/config/trainer.yaml @@ -53,7 +53,6 @@ collection: training: should: True learning_rate: 0.0001 - sampling_weights: [0.125, 0.125, 0.25, 0.5] tokenizer: batch_num_samples: 256 grad_acc_steps: 1 diff --git a/src/dataset.py b/src/dataset.py index bd0ce9c..59d9f30 100644 --- a/src/dataset.py +++ b/src/dataset.py @@ -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: diff --git a/src/trainer.py b/src/trainer.py index 24b801f..ec94572 100644 --- a/src/trainer.py +++ b/src/trainer.py @@ -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