mirror of
https://github.com/wassname/iris_bigvae.git
synced 2026-09-09 11:24:31 +08:00
Remove unused sampling_weights.
This commit is contained in:
@@ -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
|
||||
|
||||
+4
-15
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user