r""" Norm layer adaptive layer norm zero (adaLN-Zero). Parameters: embedding_dim (`int`): The size of each embedding vector. num_embeddings (`int`): The size of the embeddings dictionary.
| 98 | |
| 99 | |
| 100 | class AdaLayerNormZero(nn.Module): |
| 101 | r""" |
| 102 | Norm layer adaptive layer norm zero (adaLN-Zero). |
| 103 | |
| 104 | Parameters: |
| 105 | embedding_dim (`int`): The size of each embedding vector. |
| 106 | num_embeddings (`int`): The size of the embeddings dictionary. |
| 107 | """ |
| 108 | |
| 109 | def __init__(self, embedding_dim: int, num_embeddings: Optional[int] = None, norm_type="layer_norm", bias=True): |
| 110 | super().__init__() |
| 111 | if num_embeddings is not None: |
| 112 | self.emb = CombinedTimestepLabelEmbeddings(num_embeddings, embedding_dim) |
| 113 | else: |
| 114 | self.emb = None |
| 115 | |
| 116 | self.silu = nn.SiLU() |
| 117 | self.linear = nn.Linear(embedding_dim, 6 * embedding_dim, bias=bias) |
| 118 | if norm_type == "layer_norm": |
| 119 | self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6) |
| 120 | elif norm_type == "fp32_layer_norm": |
| 121 | self.norm = FP32LayerNorm(embedding_dim, elementwise_affine=False, bias=False) |
| 122 | else: |
| 123 | raise ValueError( |
| 124 | f"Unsupported `norm_type` ({norm_type}) provided. Supported ones are: 'layer_norm', 'fp32_layer_norm'." |
| 125 | ) |
| 126 | |
| 127 | def forward( |
| 128 | self, |
| 129 | x: torch.Tensor, |
| 130 | timestep: Optional[torch.Tensor] = None, |
| 131 | class_labels: Optional[torch.LongTensor] = None, |
| 132 | hidden_dtype: Optional[torch.dtype] = None, |
| 133 | emb: Optional[torch.Tensor] = None, |
| 134 | ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: |
| 135 | if self.emb is not None: |
| 136 | emb = self.emb(timestep, class_labels, hidden_dtype=hidden_dtype) |
| 137 | emb = self.linear(self.silu(emb)) |
| 138 | shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = emb.chunk(6, dim=1) |
| 139 | x = self.norm(x) * (1 + scale_msa[:, None]) + shift_msa[:, None] |
| 140 | return x, gate_msa, shift_mlp, scale_mlp, gate_mlp |
| 141 | |
| 142 | |
| 143 | class AdaLayerNormZeroSingle(nn.Module): |