Method
__init__
(
self,
width,
height,
hidden_size,
num_layers,
time_embed_dim,
compressed_num_frames,
qk_ln=True,
hidden_size_head=None,
elementwise_affine=True,
)
Source from the content-addressed store, hash-verified
| 573 | # * Main Transformer Layer |
| 574 | class AdaLNMixin(BaseMixin): |
| 575 | def __init__( |
| 576 | self, |
| 577 | width, |
| 578 | height, |
| 579 | hidden_size, |
| 580 | num_layers, |
| 581 | time_embed_dim, |
| 582 | compressed_num_frames, |
| 583 | qk_ln=True, |
| 584 | hidden_size_head=None, |
| 585 | elementwise_affine=True, |
| 586 | ): |
| 587 | super().__init__() |
| 588 | self.num_layers = num_layers |
| 589 | self.width = width |
| 590 | self.height = height |
| 591 | self.compressed_num_frames = compressed_num_frames |
| 592 | |
| 593 | self.adaLN_modulations = nn.ModuleList( |
| 594 | [nn.Sequential(nn.SiLU(), nn.Linear(time_embed_dim, 12 * hidden_size)) for _ in range(num_layers)] |
| 595 | ) |
| 596 | |
| 597 | self.qk_ln = qk_ln |
| 598 | if qk_ln: |
| 599 | self.query_layernorm_list = nn.ModuleList( |
| 600 | [ |
| 601 | LayerNorm(hidden_size_head, eps=1e-6, elementwise_affine=elementwise_affine) |
| 602 | for _ in range(num_layers) |
| 603 | ] |
| 604 | ) |
| 605 | self.key_layernorm_list = nn.ModuleList( |
| 606 | [ |
| 607 | LayerNorm(hidden_size_head, eps=1e-6, elementwise_affine=elementwise_affine) |
| 608 | for _ in range(num_layers) |
| 609 | ] |
| 610 | ) |
| 611 | |
| 612 | def separate_modulate( |
| 613 | self, |
Tested by
no test coverage detected