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

Class JointTransformerFinalBlock

diffsynth/models/sd3_dit.py:294–322  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

292
293
294class JointTransformerFinalBlock(torch.nn.Module):
295 def __init__(self, dim, num_attention_heads, use_rms_norm=False):
296 super().__init__()
297 self.norm1_a = AdaLayerNorm(dim)
298 self.norm1_b = AdaLayerNorm(dim, single=True)
299
300 self.attn = JointAttention(dim, dim, num_attention_heads, dim // num_attention_heads, only_out_a=True, use_rms_norm=use_rms_norm)
301
302 self.norm2_a = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
303 self.ff_a = torch.nn.Sequential(
304 torch.nn.Linear(dim, dim*4),
305 torch.nn.GELU(approximate="tanh"),
306 torch.nn.Linear(dim*4, dim)
307 )
308
309
310 def forward(self, hidden_states_a, hidden_states_b, temb):
311 norm_hidden_states_a, gate_msa_a, shift_mlp_a, scale_mlp_a, gate_mlp_a = self.norm1_a(hidden_states_a, emb=temb)
312 norm_hidden_states_b = self.norm1_b(hidden_states_b, emb=temb)
313
314 # Attention
315 attn_output_a = self.attn(norm_hidden_states_a, norm_hidden_states_b)
316
317 # Part A
318 hidden_states_a = hidden_states_a + gate_msa_a * attn_output_a
319 norm_hidden_states_a = self.norm2_a(hidden_states_a) * (1 + scale_mlp_a) + shift_mlp_a
320 hidden_states_a = hidden_states_a + gate_mlp_a * self.ff_a(norm_hidden_states_a)
321
322 return hidden_states_a, hidden_states_b
323
324
325

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected