| 109 | |
| 110 | # ET + RT + Src -> ((Role)(Role)), ETRTText2Role 使用 |
| 111 | class RoleConstraintDecoder(ConstraintDecoder): |
| 112 | def __init__(self, tokenizer, type_schema, *args, **kwargs): |
| 113 | super().__init__(tokenizer, *args, **kwargs) |
| 114 | self.tree_end = '<tree-end>' |
| 115 | self.type_schema = type_schema |
| 116 | self.type_tree = get_label_name_tree(type_schema.role_list, |
| 117 | tokenizer=self.tokenizer, |
| 118 | end_symbol=self.tree_end) |
| 119 | self.type_start = self.tokenizer.convert_tokens_to_ids([type_start])[0] |
| 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): |
| 153 | """ |
| 154 | Generate Text Span |
| 155 | :param generated: |
| 156 | :param prefix_tree: |
| 157 | :param src_sentence: |
| 158 | :param end_sequence_search_tokens: |
| 159 | :return: |
| 160 | """ |
| 161 | tree = prefix_tree |
| 162 | for index, token in enumerate(generated): |
| 163 | tree = tree[token] |
| 164 | is_tree_end = len(tree) == 1 and self.tree_end in tree |
| 165 | |
| 166 | if is_tree_end: |
| 167 | valid_token = generated_search_src_sequence( |
| 168 | generated=generated[index + 1:], |
nothing calls this directly
no outgoing calls
no test coverage detected