(self, tgt_generated)
| 381 | self.type_end = self.tokenizer.convert_tokens_to_ids([type_end])[0] |
| 382 | |
| 383 | def check_state(self, tgt_generated): |
| 384 | if tgt_generated[-1] == self.tokenizer.pad_token_id: # t5-base |
| 385 | return 'start', -1 |
| 386 | |
| 387 | special_token_set = {self.type_start, self.type_end} |
| 388 | special_index_token = list( |
| 389 | filter(lambda x: x[1] in special_token_set, list(enumerate(tgt_generated)))) |
| 390 | # print(special_index_token) |
| 391 | last_special_index, last_special_token = special_index_token[-1] |
| 392 | |
| 393 | if len(special_index_token) == 1: |
| 394 | if last_special_token != self.type_start: |
| 395 | return 'error', 0 |
| 396 | |
| 397 | bracket_position = find_bracket_position( |
| 398 | tgt_generated, _type_start=self.type_start, _type_end=self.type_end) |
| 399 | start_number, end_number = len(bracket_position[self.type_start]), len( |
| 400 | bracket_position[self.type_end]) # 计算左右括号的数量 |
| 401 | |
| 402 | if start_number == end_number: |
| 403 | return 'end_generate', -1 |
| 404 | if start_number == end_number + 1: |
| 405 | state = 'start_first_generation' |
| 406 | elif start_number == end_number + 2: |
| 407 | state = 'generate_span' |
| 408 | else: |
| 409 | state = 'error' |
| 410 | return state, last_special_index |
| 411 | |
| 412 | |
| 413 | def get_state_valid_tokens(self, src_sentence, tgt_generated): |
no test coverage detected