MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / eval

Function eval

modelzoo/dbmtl/train.py:573–597  ·  view source on GitHub ↗
(sess_config, input_hooks, model, data_init_op, steps, checkpoint_dir)

Source from the content-addressed store, hash-verified

571 print("Training completed.")
572
573def eval(sess_config, input_hooks, model, data_init_op, steps, checkpoint_dir):
574 model.is_training = False
575 hooks = []
576 hooks.extend(input_hooks)
577
578 scaffold = tf.train.Scaffold(
579 local_init_op=tf.group(tf.local_variables_initializer(), data_init_op))
580 session_creator = tf.train.ChiefSessionCreator(
581 scaffold=scaffold, checkpoint_dir=checkpoint_dir, config=sess_config)
582 writer = tf.summary.FileWriter(os.path.join(checkpoint_dir, 'eval'))
583 merged = tf.summary.merge_all()
584
585 with tf.train.MonitoredSession(session_creator=session_creator,
586 hooks=hooks) as sess:
587 for _in in range(1, steps + 1):
588 if (_in != steps):
589 sess.run([model.acc_op, model.auc_op])
590 if (_in % 1000 == 0):
591 print("Evaluation complete:[{}/{}]".format(_in, steps))
592 else:
593 eval_acc, eval_auc, events = sess.run(
594 [model.acc_op, model.auc_op, merged])
595 writer.add_summary(events, _in)
596 print("Evaluation complete:[{}/{}]".format(_in, steps))
597 print("ACC = {}\nAUC = {}".format(eval_acc, eval_auc))
598
599def main(tf_config=None, server=None):
600 # check dataset and count data set size

Callers 1

mainFunction · 0.70

Calls 7

rangeFunction · 0.50
extendMethod · 0.45
groupMethod · 0.45
joinMethod · 0.45
runMethod · 0.45
formatMethod · 0.45
add_summaryMethod · 0.45

Tested by

no test coverage detected