MCPcopy Create free account
hub / github.com/YesianRohn/TextSSR / AdaLayerNormZero

Class AdaLayerNormZero

diffusers/src/diffusers/models/normalization.py:100–140  ·  view source on GitHub ↗

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.

Source from the content-addressed store, hash-verified

98
99
100class 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
143class AdaLayerNormZeroSingle(nn.Module):

Callers 5

__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected