| 36 | |
| 37 | |
| 38 | class WanAttnProcessor2_0: |
| 39 | def __init__(self): |
| 40 | if not hasattr(F, "scaled_dot_product_attention"): |
| 41 | raise ImportError("WanAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.") |
| 42 | |
| 43 | def __call__( |
| 44 | self, |
| 45 | attn: Attention, |
| 46 | hidden_states: torch.Tensor, |
| 47 | encoder_hidden_states: Optional[torch.Tensor] = None, |
| 48 | attention_mask: Optional[torch.Tensor] = None, |
| 49 | rotary_emb: Optional[torch.Tensor] = None, |
| 50 | ) -> torch.Tensor: |
| 51 | encoder_hidden_states_img = None |
| 52 | if attn.add_k_proj is not None: |
| 53 | # 512 is the context length of the text encoder, hardcoded for now |
| 54 | image_context_length = encoder_hidden_states.shape[1] - 512 |
| 55 | encoder_hidden_states_img = encoder_hidden_states[:, :image_context_length] |
| 56 | encoder_hidden_states = encoder_hidden_states[:, image_context_length:] |
| 57 | if encoder_hidden_states is None: |
| 58 | encoder_hidden_states = hidden_states |
| 59 | |
| 60 | query = attn.to_q(hidden_states) |
| 61 | key = attn.to_k(encoder_hidden_states) |
| 62 | value = attn.to_v(encoder_hidden_states) |
| 63 | |
| 64 | if attn.norm_q is not None: |
| 65 | query = attn.norm_q(query) |
| 66 | if attn.norm_k is not None: |
| 67 | key = attn.norm_k(key) |
| 68 | |
| 69 | query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2) |
| 70 | key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2) |
| 71 | value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2) |
| 72 | |
| 73 | if rotary_emb is not None: |
| 74 | |
| 75 | def apply_rotary_emb( |
| 76 | hidden_states: torch.Tensor, |
| 77 | freqs_cos: torch.Tensor, |
| 78 | freqs_sin: torch.Tensor, |
| 79 | ): |
| 80 | x = hidden_states.view(*hidden_states.shape[:-1], -1, 2) |
| 81 | x1, x2 = x[..., 0], x[..., 1] |
| 82 | cos = freqs_cos[..., 0::2] |
| 83 | sin = freqs_sin[..., 1::2] |
| 84 | out = torch.empty_like(hidden_states) |
| 85 | out[..., 0::2] = x1 * cos - x2 * sin |
| 86 | out[..., 1::2] = x1 * sin + x2 * cos |
| 87 | return out.type_as(hidden_states) |
| 88 | |
| 89 | query = apply_rotary_emb(query, *rotary_emb) |
| 90 | key = apply_rotary_emb(key, *rotary_emb) |
| 91 | |
| 92 | # I2V task |
| 93 | hidden_states_img = None |
| 94 | if encoder_hidden_states_img is not None: |
| 95 | key_img = attn.add_k_proj(encoder_hidden_states_img) |