(self, src_sentence, tgt_generated)
| 84 | pass |
| 85 | |
| 86 | def constraint_decoding(self, src_sentence, tgt_generated): |
| 87 | if self.source_prefix_tokenized: |
| 88 | # Remove Source Prefix for Generation |
| 89 | src_sentence = src_sentence[len(self.source_prefix_tokenized):] |
| 90 | |
| 91 | if debug: |
| 92 | print("Src:", self.tokenizer.convert_ids_to_tokens(src_sentence)) |
| 93 | print("Tgt:", self.tokenizer.convert_ids_to_tokens(tgt_generated)) |
| 94 | |
| 95 | valid_token_ids = self.get_state_valid_tokens( |
| 96 | src_sentence.tolist(), |
| 97 | tgt_generated.tolist() |
| 98 | ) |
| 99 | |
| 100 | if debug: |
| 101 | print('========================================') |
| 102 | print('valid tokens:', self.tokenizer.convert_ids_to_tokens( |
| 103 | valid_token_ids), valid_token_ids) |
| 104 | if debug_step: |
| 105 | input() |
| 106 | |
| 107 | # return self.tokenizer.convert_tokens_to_ids(valid_tokens) |
| 108 | return valid_token_ids |
| 109 | |
| 110 | # ET + RT + Src -> ((Role)(Role)), ETRTText2Role 使用 |
| 111 | class RoleConstraintDecoder(ConstraintDecoder): |
no test coverage detected