MCPcopy Create free account
hub / github.com/GanjinZero/BioBART / create_training_instance

Method create_training_instance

pretrain_src/dataloader.py:109–131  ·  view source on GitHub ↗
(self, instance: TokenInstance)

Source from the content-addressed store, hash-verified

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

Callers 1

__getitem__Method · 0.95

Calls 4

create_noised_inputMethod · 0.95
padding_to_maxlengthFunction · 0.85
map_to_torchFunction · 0.85
get_valuesMethod · 0.80

Tested by

no test coverage detected