| 128 | class AdaLayerNorm(Module): |
| 129 | |
| 130 | def __init__(self, |
| 131 | embedding_dim: int, |
| 132 | num_embeddings: Optional[int] = None, |
| 133 | output_dim: Optional[int] = None, |
| 134 | norm_elementwise_affine: bool = False, |
| 135 | norm_eps: float = 1e-5, |
| 136 | chunk_dim: int = 0, |
| 137 | mapping=Mapping(), |
| 138 | dtype=None): |
| 139 | super().__init__() |
| 140 | self.chunk_dim = chunk_dim |
| 141 | output_dim = output_dim or embedding_dim * 2 |
| 142 | if num_embeddings is not None: |
| 143 | self.emb = Embedding(num_embeddings, embedding_dim, dtype=dtype) |
| 144 | else: |
| 145 | self.emb = None |
| 146 | self.silu = ACT2FN['silu'] |
| 147 | self.linear = Linear(embedding_dim, |
| 148 | output_dim, |
| 149 | tp_group=mapping.tp_group, |
| 150 | tp_size=mapping.tp_size, |
| 151 | dtype=dtype) |
| 152 | self.norm = LayerNorm(output_dim // 2, |
| 153 | eps=norm_eps, |
| 154 | elementwise_affine=norm_elementwise_affine, |
| 155 | dtype=dtype) |
| 156 | |
| 157 | def forward(self, |
| 158 | x: Tensor, |