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

Class TriConstraintDecoder

extraction/extract_constraint.py:372–473  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

370
371# ET + Src -> ((Tri)(Tri)), ETText2Tri 使用
372class 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)]

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected