MCPcopy Create free account
hub / github.com/Monalissaa/DisenDiff / custom_forward

Method custom_forward

src/custom_modules.py:257–280  ·  view source on GitHub ↗

r""" Returns:

(self, hidden_states, input_ids)

Source from the content-addressed store, hash-verified

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()

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected