| 251 | |
| 252 | # Src -> ((ET)(ET)), Text2ET 使用 |
| 253 | class 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: |
nothing calls this directly
no outgoing calls
no test coverage detected