(self, distribution_strategy, flag_mode)
| 220 | flag_mode=['train'], |
| 221 | )) |
| 222 | def test_recovery(self, distribution_strategy, flag_mode): |
| 223 | loss_threshold = 1.0 |
| 224 | model_dir = self.get_temp_dir() |
| 225 | flags_dict = dict( |
| 226 | experiment='mock', |
| 227 | mode=flag_mode, |
| 228 | model_dir=model_dir, |
| 229 | params_override=json.dumps(self._test_config)) |
| 230 | with flagsaver.flagsaver(**flags_dict): |
| 231 | params = train_utils.parse_configuration(flags.FLAGS) |
| 232 | params.trainer.loss_upper_bound = loss_threshold |
| 233 | params.trainer.recovery_max_trials = 1 |
| 234 | train_utils.serialize_config(params, model_dir) |
| 235 | with distribution_strategy.scope(): |
| 236 | task = task_factory.get_task(params.task, logging_dir=model_dir) |
| 237 | |
| 238 | # Saves a checkpoint for reference. |
| 239 | model = task.build_model() |
| 240 | checkpoint = tf.train.Checkpoint(model=model) |
| 241 | checkpoint_manager = tf.train.CheckpointManager( |
| 242 | checkpoint, self.get_temp_dir(), max_to_keep=2) |
| 243 | checkpoint_manager.save() |
| 244 | before_weights = model.get_weights() |
| 245 | |
| 246 | def build_losses(labels, model_outputs, aux_losses=None): |
| 247 | del labels, model_outputs |
| 248 | return tf.constant([loss_threshold], tf.float32) + aux_losses |
| 249 | |
| 250 | task.build_losses = build_losses |
| 251 | |
| 252 | model, _ = train_lib.OrbitExperimentRunner( |
| 253 | distribution_strategy=distribution_strategy, |
| 254 | task=task, |
| 255 | mode=flag_mode, |
| 256 | params=params, |
| 257 | model_dir=model_dir).run() |
| 258 | after_weights = model.get_weights() |
| 259 | for left, right in zip(before_weights, after_weights): |
| 260 | self.assertAllEqual(left, right) |
| 261 | |
| 262 | def test_parse_configuration(self): |
| 263 | model_dir = self.get_temp_dir() |
nothing calls this directly
no test coverage detected