(self, tgt_generated)
| 120 | self.type_end = self.tokenizer.convert_tokens_to_ids([type_end])[0] |
| 121 | |
| 122 | def check_state(self, tgt_generated): |
| 123 | if tgt_generated[-1] == self.tokenizer.pad_token_id: # t5-base |
| 124 | return 'start', -1 |
| 125 | |
| 126 | special_token_set = {self.type_start, self.type_end} |
| 127 | special_index_token = list( |
| 128 | filter(lambda x: x[1] in special_token_set, list(enumerate(tgt_generated)))) |
| 129 | # print(special_index_token) |
| 130 | last_special_index, last_special_token = special_index_token[-1] |
| 131 | |
| 132 | if len(special_index_token) == 1: |
| 133 | if last_special_token != self.type_start: |
| 134 | return 'error', 0 |
| 135 | |
| 136 | bracket_position = find_bracket_position( |
| 137 | tgt_generated, _type_start=self.type_start, _type_end=self.type_end) |
| 138 | start_number, end_number = len(bracket_position[self.type_start]), len( |
| 139 | bracket_position[self.type_end]) # 计算左右括号的数量 |
| 140 | |
| 141 | if start_number == end_number: |
| 142 | return 'end_generate', -1 |
| 143 | if start_number == end_number + 1: |
| 144 | state = 'start_first_generation' |
| 145 | elif start_number == end_number + 2: |
| 146 | state = 'generate_span' |
| 147 | else: |
| 148 | state = 'error' |
| 149 | return state, last_special_index |
| 150 | |
| 151 | def search_prefix_tree_and_sequence(self, generated: List[str], prefix_tree: Dict, src_sentence: List[str], |
| 152 | end_sequence_search_tokens: List[str] = None): |
no test coverage detected