(self, tgt_generated)
| 261 | self.type_end = self.tokenizer.convert_tokens_to_ids([type_end])[0] |
| 262 | |
| 263 | def check_state(self, tgt_generated): |
| 264 | if tgt_generated[-1] == self.tokenizer.pad_token_id: # t5-base |
| 265 | return 'start', -1 |
| 266 | |
| 267 | special_token_set = {self.type_start, self.type_end} |
| 268 | special_index_token = list( |
| 269 | filter(lambda x: x[1] in special_token_set, list(enumerate(tgt_generated)))) |
| 270 | # print(special_index_token) |
| 271 | last_special_index, last_special_token = special_index_token[-1] |
| 272 | |
| 273 | if len(special_index_token) == 1: |
| 274 | if last_special_token != self.type_start: |
| 275 | return 'error', 0 |
| 276 | |
| 277 | bracket_position = find_bracket_position( |
| 278 | tgt_generated, _type_start=self.type_start, _type_end=self.type_end) |
| 279 | start_number, end_number = len(bracket_position[self.type_start]), len( |
| 280 | bracket_position[self.type_end]) # 计算左右括号的数量 |
| 281 | |
| 282 | if start_number == end_number: |
| 283 | return 'end_generate', -1 |
| 284 | if start_number == end_number + 1: |
| 285 | state = 'start_first_generation' |
| 286 | elif start_number == end_number + 2: |
| 287 | state = 'generate_span' |
| 288 | else: |
| 289 | state = 'error' |
| 290 | return state, last_special_index |
| 291 | |
| 292 | def search_prefix_tree(self, generated: List[str], prefix_tree: Dict, |
| 293 | end_sequence_search_tokens: List[str] = None): |
no test coverage detected