| 84 | self._visualize_attention(model, args.beam_size, args.output_history, sliced_data, args.res1, args.res2, args.res3, output) |
| 85 | |
| 86 | def _infer_one(self, model, data_item, preproc_item, beam_size, output_history=False, use_heuristic=True): |
| 87 | if use_heuristic: |
| 88 | # TODO: from_cond should be true from non-bert model |
| 89 | beams = spider_beam_search.beam_search_with_heuristics( |
| 90 | model, data_item, preproc_item, beam_size=beam_size, max_steps=1000, from_cond=False) |
| 91 | else: |
| 92 | beams = beam_search.beam_search( |
| 93 | model, data_item, preproc_item, beam_size=beam_size, max_steps=1000) |
| 94 | decoded = [] |
| 95 | for beam in beams: |
| 96 | model_output, inferred_code = beam.inference_state.finalize() |
| 97 | |
| 98 | decoded.append({ |
| 99 | 'orig_question': data_item.orig["question"], |
| 100 | 'model_output': model_output, |
| 101 | 'inferred_code': inferred_code, |
| 102 | 'score': beam.score, |
| 103 | **({ |
| 104 | 'choice_history': beam.choice_history, |
| 105 | 'score_history': beam.score_history, |
| 106 | } if output_history else {})}) |
| 107 | return decoded |
| 108 | |
| 109 | def _inner_infer(self, model, beam_size, output_history, sliced_orig_data, sliced_preproc_data, output, use_heuristic=False): |
| 110 | for i, (orig_item, preproc_item) in enumerate( |