(self, instance: TokenInstance)
| 107 | return self.create_training_instance(instance) |
| 108 | |
| 109 | def create_training_instance(self, instance: TokenInstance): |
| 110 | |
| 111 | token_x, token_y, is_rotate = instance.get_values() |
| 112 | |
| 113 | x = [] |
| 114 | y = [] |
| 115 | |
| 116 | y = token_y + ['</s>'] |
| 117 | |
| 118 | # Get Masked LM predictions |
| 119 | noised_tokens, number_masked_tokens = self.create_noised_input(token_x) |
| 120 | # noised_tokens = token_x |
| 121 | |
| 122 | x.append('<s>') |
| 123 | x = x + noised_tokens |
| 124 | x.append('</s>') |
| 125 | |
| 126 | input_ids, attn_mask = padding_to_maxlength(self.tokenizer.convert_tokens_to_ids(x), self.max_seq_length) |
| 127 | labels, decoder_attn_mask = padding_to_maxlength(self.tokenizer.convert_tokens_to_ids(y), self.max_seq_length) |
| 128 | # input_ids, attn_mask = padding_to_maxlength(self.tokenizer.encode(x, is_pretokenized=True, add_special_tokens = False).ids[:self.max_seq_length], self.max_seq_length) |
| 129 | # labels, decoder_attn_mask = padding_to_maxlength(self.tokenizer.encode(y, is_pretokenized=True, add_special_tokens = False).ids[:self.max_seq_length], self.max_seq_length) |
| 130 | |
| 131 | return [map_to_torch(input_ids), map_to_torch(labels), map_to_torch(attn_mask), map_to_torch(decoder_attn_mask)] |
| 132 | |
| 133 | def create_noised_input(self, tokens_x): |
| 134 | masked_number = 0 |
no test coverage detected