| 370 | |
| 371 | # ET + Src -> ((Tri)(Tri)), ETText2Tri 使用 |
| 372 | class TriConstraintDecoder(ConstraintDecoder): |
| 373 | def __init__(self, tokenizer, type_schema, *args, **kwargs): |
| 374 | super().__init__(tokenizer, *args, **kwargs) |
| 375 | self.tree_end = '<tree-end>' |
| 376 | self.type_schema = type_schema |
| 377 | self.type_tree = get_label_name_tree(type_schema.role_list, |
| 378 | tokenizer=self.tokenizer, |
| 379 | end_symbol=self.tree_end) |
| 380 | self.type_start = self.tokenizer.convert_tokens_to_ids([type_start])[0] |
| 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): |
| 414 | """ |
| 415 | |
| 416 | :param src_sentence: ET </s> src </s> |
| 417 | :param tgt_generated: |
| 418 | :return: |
| 419 | List[str], valid token list |
| 420 | """ |
| 421 | old_src = src_sentence |
| 422 | if self.tokenizer.eos_token_id in src_sentence: |
| 423 | if src_sentence.count(self.tokenizer.eos_token_id) > 1: # 有新增的</s> |
| 424 | first_index = src_sentence.index(self.tokenizer.eos_token_id) # index函数会定位第一个出现的位置 |
| 425 | second_index = first_index + 1 + src_sentence[first_index + 1:].index(self.tokenizer.eos_token_id) # 注意要加上 first_index + 1的偏移 |
| 426 | src_sentence = src_sentence[first_index + 1: second_index] # 输入端 原句 src |
| 427 | |
| 428 | else: |
| 429 | src_sentence = src_sentence[:src_sentence.index(self.tokenizer.eos_token_id)] |
nothing calls this directly
no outgoing calls
no test coverage detected