MCPcopy Create free account
hub / github.com/UVA-Computer-Vision-Lab/FrameINO / WanAttnProcessor2_0

Class WanAttnProcessor2_0

architecture/transformer_wan.py:38–119  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

36
37
38class 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)

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected