MCPcopy Create free account
hub / github.com/OpenImagingLab/FlashVSR / TemporalAttentionBlock

Class TemporalAttentionBlock

diffsynth/models/svd_unet.py:138–214  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

136
137
138class TemporalAttentionBlock(torch.nn.Module):
139
140 def __init__(self, num_attention_heads, attention_head_dim, in_channels, cross_attention_dim=None, add_positional_conv=None):
141 super().__init__()
142
143 self.positional_embedding_proj = torch.nn.Sequential(
144 torch.nn.Linear(in_channels, in_channels * 4),
145 torch.nn.SiLU(),
146 torch.nn.Linear(in_channels * 4, in_channels)
147 )
148 if add_positional_conv is not None:
149 self.positional_embedding = TrainableTemporalTimesteps(in_channels, True, 0, add_positional_conv)
150 self.positional_conv = torch.nn.Conv3d(in_channels, in_channels, kernel_size=3, padding=1, padding_mode="reflect")
151 else:
152 self.positional_embedding = TemporalTimesteps(in_channels, True, 0)
153 self.positional_conv = None
154
155 self.norm_in = torch.nn.LayerNorm(in_channels)
156 self.act_fn_in = GEGLU(in_channels, in_channels * 4)
157 self.ff_in = torch.nn.Linear(in_channels * 4, in_channels)
158
159 self.norm1 = torch.nn.LayerNorm(in_channels)
160 self.attn1 = Attention(
161 q_dim=in_channels,
162 num_heads=num_attention_heads,
163 head_dim=attention_head_dim,
164 bias_out=True
165 )
166
167 self.norm2 = torch.nn.LayerNorm(in_channels)
168 self.attn2 = Attention(
169 q_dim=in_channels,
170 kv_dim=cross_attention_dim,
171 num_heads=num_attention_heads,
172 head_dim=attention_head_dim,
173 bias_out=True
174 )
175
176 self.norm_out = torch.nn.LayerNorm(in_channels)
177 self.act_fn_out = GEGLU(in_channels, in_channels * 4)
178 self.ff_out = torch.nn.Linear(in_channels * 4, in_channels)
179
180 def forward(self, hidden_states, time_emb, text_emb, res_stack, **kwargs):
181
182 batch, inner_dim, height, width = hidden_states.shape
183 pos_emb = torch.arange(batch)
184 pos_emb = self.positional_embedding(pos_emb).to(dtype=hidden_states.dtype, device=hidden_states.device)
185 pos_emb = self.positional_embedding_proj(pos_emb)
186
187 hidden_states = rearrange(hidden_states, "T C H W -> 1 C T H W") + rearrange(pos_emb, "T C -> 1 C T 1 1")
188 if self.positional_conv is not None:
189 hidden_states = self.positional_conv(hidden_states)
190 hidden_states = rearrange(hidden_states[0], "C T H W -> (H W) T C")
191
192 residual = hidden_states
193 hidden_states = self.norm_in(hidden_states)
194 hidden_states = self.act_fn_in(hidden_states)
195 hidden_states = self.ff_in(hidden_states)

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected