MCPcopy Create free account
hub / github.com/JonasGeiping/cramming / ScriptableLMForPreTraining

Class ScriptableLMForPreTraining

cramming/architectures/scriptable_bert.py:138–219  ·  view source on GitHub ↗

Definitely can represent BERT, but also a lot of other things. To be used for MLM schemes.

Source from the content-addressed store, hash-verified

136
137
138class ScriptableLMForPreTraining(PreTrainedModel):
139 """Definitely can represent BERT, but also a lot of other things. To be used for MLM schemes."""
140
141 config_class = crammedBertConfig
142
143 def __init__(self, config):
144 super().__init__(config)
145 self.cfg = OmegaConf.create(config.arch) # this could be nicer ...
146 self.encoder = ScriptableLM(config)
147 if not self.cfg.skip_head_transform:
148 self.prediction_head = PredictionHeadComponent(self.cfg)
149 else:
150 self.prediction_head = torch.nn.Linear(
151 self.cfg.hidden_size,
152 self.cfg.embedding.embedding_dim,
153 bias=self.cfg.use_bias,
154 )
155
156 if self.cfg.loss == "szegedy":
157 self.decoder = torch.nn.Identity()
158 else:
159 if self.cfg.tie_weights:
160 self.decoder = torch.nn.Linear(self.cfg.embedding.embedding_dim, self.cfg.embedding.vocab_size, bias=self.cfg.decoder_bias)
161 self.decoder.weight = self.encoder.embedding.word_embedding.weight
162 else:
163 self.decoder = torch.nn.Linear(self.cfg.hidden_size, self.cfg.embedding.vocab_size, bias=self.cfg.decoder_bias)
164
165 self.loss_fn = _get_loss_fn(self.cfg.loss, z_loss_factor=self.cfg.z_loss_factor, embedding=self.encoder.embedding.word_embedding)
166 self.sparse_prediction = self.cfg.sparse_prediction
167 self.vocab_size = self.cfg.embedding.vocab_size
168
169 self._init_weights()
170
171 def _init_weights(self, *args, **kwargs):
172 for name, module in self.named_modules():
173 _init_module(
174 name,
175 module,
176 self.cfg.init.type,
177 self.cfg.init.std,
178 self.cfg.hidden_size,
179 self.cfg.num_transformer_layers,
180 )
181
182 def forward(
183 self,
184 input_ids,
185 attention_mask: Optional[torch.Tensor] = None,
186 labels: Optional[torch.Tensor] = None,
187 token_type_ids: Optional[torch.Tensor] = None,
188 ):
189 outputs = self.encoder(input_ids, attention_mask)
190 outputs = outputs.view(-1, outputs.shape[-1])
191
192 if self.sparse_prediction:
193 masked_lm_loss = self._forward_dynamic(outputs, labels)
194 else:
195 outputs = self.decoder(self.prediction_head(outputs))

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected