(self, hidden_states, prompt_emb, time_emb, image_rotary_emb)
| 78 | |
| 79 | |
| 80 | def forward(self, hidden_states, prompt_emb, time_emb, image_rotary_emb): |
| 81 | # Attention |
| 82 | norm_hidden_states, norm_encoder_hidden_states, gate_a, gate_b = self.norm1( |
| 83 | hidden_states, prompt_emb, time_emb |
| 84 | ) |
| 85 | attention_io = torch.cat([norm_encoder_hidden_states, norm_hidden_states], dim=1) |
| 86 | attention_io = self.attn1( |
| 87 | attention_io, |
| 88 | qkv_preprocessor=lambda q, k, v: self.process_qkv(q, k, v, image_rotary_emb, prompt_emb.shape[1]) |
| 89 | ) |
| 90 | |
| 91 | hidden_states = hidden_states + gate_a * attention_io[:, prompt_emb.shape[1]:] |
| 92 | prompt_emb = prompt_emb + gate_b * attention_io[:, :prompt_emb.shape[1]] |
| 93 | |
| 94 | # Feed forward |
| 95 | norm_hidden_states, norm_encoder_hidden_states, gate_a, gate_b = self.norm2( |
| 96 | hidden_states, prompt_emb, time_emb |
| 97 | ) |
| 98 | ff_io = torch.cat([norm_encoder_hidden_states, norm_hidden_states], dim=1) |
| 99 | ff_io = self.ff(ff_io) |
| 100 | |
| 101 | hidden_states = hidden_states + gate_a * ff_io[:, prompt_emb.shape[1]:] |
| 102 | prompt_emb = prompt_emb + gate_b * ff_io[:, :prompt_emb.shape[1]] |
| 103 | |
| 104 | return hidden_states, prompt_emb |
| 105 | |
| 106 | |
| 107 |
nothing calls this directly
no test coverage detected