|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| """Multitask trainer that interleaves each task's train step."""
|
| from typing import Union
|
| import gin
|
| import orbit
|
| import tensorflow as tf, tf_keras
|
| from official.modeling.multitask import base_model
|
| from official.modeling.multitask import base_trainer
|
| from official.modeling.multitask import multitask
|
| from official.modeling.multitask import task_sampler as sampler
|
|
|
|
|
| @gin.configurable
|
| class MultiTaskInterleavingTrainer(base_trainer.MultiTaskBaseTrainer):
|
| """MultiTask trainer that interleaves task update."""
|
|
|
| def __init__(self,
|
| multi_task: multitask.MultiTask,
|
| multi_task_model: Union[tf_keras.Model,
|
| base_model.MultiTaskBaseModel],
|
| optimizer: Union[tf.optimizers.Optimizer,
|
| tf_keras.optimizers.experimental.Optimizer,
|
| tf_keras.optimizers.legacy.Optimizer],
|
| task_sampler: sampler.TaskSampler,
|
| trainer_options=None):
|
| super().__init__(
|
| multi_task=multi_task,
|
| multi_task_model=multi_task_model,
|
| optimizer=optimizer,
|
| trainer_options=trainer_options)
|
| self._task_sampler = task_sampler
|
|
|
|
|
| def _get_task_step(task_name, task):
|
|
|
| def step_fn(inputs):
|
| if isinstance(self.multi_task_model, base_model.MultiTaskBaseModel):
|
| task_model = self.multi_task_model.sub_tasks[task_name]
|
| else:
|
| task_model = self.multi_task_model
|
| task_logs = task.train_step(
|
| inputs,
|
| model=task_model,
|
| optimizer=self.optimizer,
|
| metrics=self.training_metrics[task_name])
|
| self.training_losses[task_name].update_state(task_logs[task.loss])
|
|
|
| return step_fn
|
|
|
| self._task_train_step_map = {
|
| name: _get_task_step(name, task)
|
| for name, task in self.multi_task.tasks.items()
|
| }
|
|
|
|
|
|
|
| self._task_step_counters = {
|
| name: orbit.utils.create_global_step() for name in self.multi_task.tasks
|
| }
|
|
|
|
|
|
|
|
|
| if isinstance(optimizer, tf_keras.optimizers.experimental.Optimizer):
|
| multi_task_model.build()
|
| optimizer.build(multi_task_model.trainable_variables)
|
|
|
| def task_step_counter(self, name):
|
| return self._task_step_counters[name]
|
|
|
| def train_step(self, iterator_map):
|
|
|
| rn = tf.random.stateless_uniform(shape=[], seed=(0, self.global_step))
|
| cumulative_sample_distribution = self._task_sampler.task_cumulative_distribution(
|
| self.global_step)
|
|
|
| cumulative_sample_distribution = tf.concat(
|
| [tf.constant([0.0], dtype=tf.float32), cumulative_sample_distribution],
|
| axis=0)
|
|
|
| for idx, (name, _) in enumerate(self.multi_task.tasks.items()):
|
| begin = cumulative_sample_distribution[idx]
|
| end = cumulative_sample_distribution[idx + 1]
|
| if rn >= begin and rn < end:
|
| self._strategy.run(
|
| self._task_train_step_map[name], args=(next(iterator_map[name]),))
|
| self.global_step.assign_add(1)
|
| self.task_step_counter(name).assign_add(1)
|
|
|
| def train_loop_end(self):
|
| """Record loss and metric values per task."""
|
| result = super().train_loop_end()
|
|
|
|
|
|
|
| if 'total_loss' in result:
|
| result.pop('total_loss')
|
| return result
|
|
|