Definitely can represent BERT, but also a lot of other things. To be used for MLM schemes.
| 136 | |
| 137 | |
| 138 | class 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)) |
no outgoing calls
no test coverage detected