MCPcopy Create free account
hub / github.com/RingBDStack/GDAP / RoleConstraintDecoder

Class RoleConstraintDecoder

extraction/extract_constraint.py:111–250  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

109
110# ET + RT + Src -> ((Role)(Role)), ETRTText2Role 使用
111class 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:],

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected