Reformats BERT input features.
(self, bert_inputs, fmt)
| 149 | } |
| 150 | |
| 151 | def _reformat_bert_inputs(self, bert_inputs, fmt): |
| 152 | """Reformats BERT input features.""" |
| 153 | if fmt == 'bert_inputs': |
| 154 | return bert_inputs # No reformatting necessary. |
| 155 | elif fmt == 'dict': |
| 156 | # This is the format expected by the original BERT code. |
| 157 | return { |
| 158 | 'input_ids': bert_inputs.token_ids, |
| 159 | 'input_mask': bert_inputs.mask, |
| 160 | 'segment_ids': bert_inputs.segment_ids, |
| 161 | } |
| 162 | else: |
| 163 | raise ValueError('Invalid format: {}'.format(fmt)) |
| 164 | |
| 165 | @profile.profiled_function |
| 166 | def featurize_query(self, query, fmt='dict'): |
no outgoing calls
no test coverage detected