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

Method __init__

diffsynth/models/sd3_dit.py:327–337  ·  view source on GitHub ↗
(self, embed_dim=1536, num_layers=24, use_rms_norm=False, num_dual_blocks=0, pos_embed_max_size=192)

Source from the content-addressed store, hash-verified

325
326class SD3DiT(torch.nn.Module):
327 def __init__(self, embed_dim=1536, num_layers=24, use_rms_norm=False, num_dual_blocks=0, pos_embed_max_size=192):
328 super().__init__()
329 self.pos_embedder = PatchEmbed(patch_size=2, in_channels=16, embed_dim=embed_dim, pos_embed_max_size=pos_embed_max_size)
330 self.time_embedder = TimestepEmbeddings(256, embed_dim)
331 self.pooled_text_embedder = torch.nn.Sequential(torch.nn.Linear(2048, embed_dim), torch.nn.SiLU(), torch.nn.Linear(embed_dim, embed_dim))
332 self.context_embedder = torch.nn.Linear(4096, embed_dim)
333 self.blocks = torch.nn.ModuleList([JointTransformerBlock(embed_dim, embed_dim//64, use_rms_norm=use_rms_norm, dual=True) for _ in range(num_dual_blocks)]
334 + [JointTransformerBlock(embed_dim, embed_dim//64, use_rms_norm=use_rms_norm) for _ in range(num_layers-1-num_dual_blocks)]
335 + [JointTransformerFinalBlock(embed_dim, embed_dim//64, use_rms_norm=use_rms_norm)])
336 self.norm_out = AdaLayerNorm(embed_dim, single=True)
337 self.proj_out = torch.nn.Linear(embed_dim, 64)
338
339 def tiled_forward(self, hidden_states, timestep, prompt_emb, pooled_prompt_emb, tile_size=128, tile_stride=64):
340 # Due to the global positional embedding, we cannot implement layer-wise tiled forward.

Callers 9

__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45

Calls 5

TimestepEmbeddingsClass · 0.85
AdaLayerNormClass · 0.85
PatchEmbedClass · 0.70

Tested by

no test coverage detected