MCPcopy Create free account
hub / github.com/tensorflow/models / joint_train_step

Method joint_train_step

official/modeling/multitask/multitask.py:104–149  ·  view source on GitHub ↗

The joint train step. Args: task_inputs: a dictionary of task names and per-task features. multi_task_model: a MultiTaskBaseModel instance. optimizer: a tf.optimizers.Optimizer. task_metrics: a dictionary of task names and per-task metrics. **kwargs: other argument

(self, task_inputs,
                       multi_task_model: base_model.MultiTaskBaseModel,
                       optimizer: tf_keras.optimizers.Optimizer, task_metrics,
                       **kwargs)

Source from the content-addressed store, hash-verified

102 dp_config=dp_config)
103
104 def joint_train_step(self, task_inputs,
105 multi_task_model: base_model.MultiTaskBaseModel,
106 optimizer: tf_keras.optimizers.Optimizer, task_metrics,
107 **kwargs):
108 """The joint train step.
109
110 Args:
111 task_inputs: a dictionary of task names and per-task features.
112 multi_task_model: a MultiTaskBaseModel instance.
113 optimizer: a tf.optimizers.Optimizer.
114 task_metrics: a dictionary of task names and per-task metrics.
115 **kwargs: other arguments to pass through.
116
117 Returns:
118 A dictionary of losses, inculding per-task losses and their weighted sum.
119 """
120 losses = {}
121 with tf.GradientTape() as tape:
122 total_loss = 0.0
123 for name, model in multi_task_model.sub_tasks.items():
124 inputs = task_inputs[name]
125 if isinstance(inputs, tuple) and len(inputs) == 2:
126 features, labels = inputs
127 elif isinstance(inputs, dict):
128 features, labels = inputs, inputs
129 else:
130 raise ValueError("The iterator output is neither a tuple nor a "
131 "dictionary. It is not implemented to support "
132 "such outputs.")
133 outputs = model(features, training=True)
134 task_loss = self.tasks[name].build_losses(labels, outputs)
135 task_weight = self.task_weight(name)
136 total_loss += task_weight * task_loss
137 losses[name] = task_loss
138 self.tasks[name].process_metrics(task_metrics[name], labels, outputs,
139 **kwargs)
140
141 # Scales loss as the default gradients allreduce performs sum inside
142 # the optimizer.
143 scaled_loss = total_loss / tf.distribute.get_strategy(
144 ).num_replicas_in_sync
145 tvars = multi_task_model.trainable_variables
146 grads = tape.gradient(scaled_loss, tvars)
147 optimizer.apply_gradients(list(zip(grads, tvars)))
148 losses["total_loss"] = total_loss
149 return losses

Callers 1

step_fnMethod · 0.80

Calls 4

task_weightMethod · 0.95
build_lossesMethod · 0.45
process_metricsMethod · 0.45
apply_gradientsMethod · 0.45

Tested by

no test coverage detected