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

Class ETConstraintDecoder

extraction/extract_constraint.py:253–369  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

251
252# Src -> ((ET)(ET)), Text2ET 使用
253class ETConstraintDecoder(ConstraintDecoder):
254 def __init__(self, tokenizer, type_schema, *args, **kwargs):
255 super().__init__(tokenizer, *args, **kwargs)
256 self.tree_end = '<tree-end>'
257 self.type_tree = get_label_name_tree(type_schema.type_list,
258 tokenizer=self.tokenizer,
259 end_symbol=self.tree_end)
260 self.type_start = self.tokenizer.convert_tokens_to_ids([type_start])[0]
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):
294 """
295 Generate Text Span
296 :param generated:
297 :param prefix_tree:
298 :param src_sentence:
299 :param end_sequence_search_tokens:
300 :return:
301 """
302 tree = prefix_tree
303 for index, token in enumerate(generated):
304 tree = tree[token]
305 is_tree_end = len(tree) == 1 and self.tree_end in tree
306
307 if is_tree_end:
308 return end_sequence_search_tokens
309
310 if self.tree_end in tree:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected