(eval_fn, input_fn, decode_fn, path, config, device_list)
| 157 | |
| 158 | |
| 159 | def _evaluate(eval_fn, input_fn, decode_fn, path, config, device_list): |
| 160 | graph = tf.Graph() |
| 161 | with graph.as_default(): |
| 162 | features = input_fn() |
| 163 | refs = features["references"] |
| 164 | placeholders = [] |
| 165 | for i in range(len(device_list)): |
| 166 | placeholders.append({ |
| 167 | "source": tf.placeholder(tf.int32, [None, None], |
| 168 | "source_%d" % i), |
| 169 | "source_length": tf.placeholder(tf.int32, [None], |
| 170 | "source_length_%d" % i) |
| 171 | }) |
| 172 | predictions = parallel.data_parallelism( |
| 173 | device_list, eval_fn, placeholders) |
| 174 | predictions = [pred[0][:, 0, :] for pred in predictions] |
| 175 | |
| 176 | all_refs = [[] for _ in range(len(refs))] |
| 177 | all_outputs = [] |
| 178 | |
| 179 | sess_creator = tf.train.ChiefSessionCreator( |
| 180 | checkpoint_dir=path, |
| 181 | config=config |
| 182 | ) |
| 183 | |
| 184 | with tf.train.MonitoredSession(session_creator=sess_creator) as sess: |
| 185 | while not sess.should_stop(): |
| 186 | feats = sess.run(features) |
| 187 | inp_feats = { |
| 188 | "source": feats["source"], |
| 189 | "source_length": feats["source_length"] |
| 190 | } |
| 191 | op, feed_dict = _shard_features(inp_feats, placeholders, |
| 192 | predictions) |
| 193 | # A list of numpy array with shape: [batch, len] |
| 194 | outputs = sess.run(op, feed_dict=feed_dict) |
| 195 | |
| 196 | for shard in outputs: |
| 197 | all_outputs.extend(shard.tolist()) |
| 198 | |
| 199 | # shape: ([batch, len], ..., [batch, len]) |
| 200 | references = [item.tolist() for item in feats["references"]] |
| 201 | |
| 202 | for i in range(len(refs)): |
| 203 | all_refs[i].extend(references[i]) |
| 204 | |
| 205 | decoded_symbols = decode_fn(all_outputs) |
| 206 | |
| 207 | for i, l in enumerate(decoded_symbols): |
| 208 | decoded_symbols[i] = " ".join(l).replace("@@ ", "").split() |
| 209 | |
| 210 | decoded_refs = [decode_fn(refs) for refs in all_refs] |
| 211 | decoded_refs = [list(x) for x in zip(*decoded_refs)] |
| 212 | |
| 213 | return bleu.bleu(decoded_symbols, decoded_refs) |
| 214 | |
| 215 | |
| 216 | class EvaluationHook(tf.train.SessionRunHook): |
no test coverage detected