(self, embed_dim=1536, num_layers=24, use_rms_norm=False, num_dual_blocks=0, pos_embed_max_size=192)
| 325 | |
| 326 | class 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. |
no test coverage detected