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

Method test_recovery

official/core/train_lib_test.py:222–260  ·  view source on GitHub ↗
(self, distribution_strategy, flag_mode)

Source from the content-addressed store, hash-verified

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()

Callers

nothing calls this directly

Calls 5

get_temp_dirMethod · 0.80
saveMethod · 0.80
build_modelMethod · 0.45
get_weightsMethod · 0.45
runMethod · 0.45

Tested by

no test coverage detected