r""" Returns:
(self, hidden_states, input_ids)
| 255 | token_embeds[self.modifier_token_id[-3]] = torch.nn.Parameter(token_embeds[43514], requires_grad=True) |
| 256 | |
| 257 | def custom_forward(self, hidden_states, input_ids): |
| 258 | r""" |
| 259 | Returns: |
| 260 | """ |
| 261 | input_shape = hidden_states.size() |
| 262 | bsz, seq_len = input_shape[:2] |
| 263 | if version.parse(transformers.__version__) >= version.parse('4.21'): |
| 264 | causal_attention_mask = self.transformer.text_model._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to( |
| 265 | hidden_states.device |
| 266 | ) |
| 267 | else: |
| 268 | causal_attention_mask = self.transformer.text_model._build_causal_attention_mask(bsz, seq_len).to( |
| 269 | hidden_states.device |
| 270 | ) |
| 271 | |
| 272 | encoder_outputs = self.transformer.text_model.encoder( |
| 273 | inputs_embeds=hidden_states, |
| 274 | causal_attention_mask=causal_attention_mask, |
| 275 | ) |
| 276 | |
| 277 | last_hidden_state = encoder_outputs[0] |
| 278 | last_hidden_state = self.transformer.text_model.final_layer_norm(last_hidden_state) |
| 279 | |
| 280 | return last_hidden_state |
| 281 | |
| 282 | def freeze(self): |
| 283 | self.transformer = self.transformer.eval() |