@@ -197,7 +197,12 @@ def __init__(
197197 num_features = embedding_dim , epsilon = eps , use_bias = False , use_scale = False , dtype = dtype , rngs = rngs
198198 )
199199 self .linear = nnx .Linear (
200- in_features = embedding_dim , out_features = embedding_dim * 2 , use_bias = True , dtype = dtype , param_dtype = weights_dtype , rngs = rngs
200+ in_features = embedding_dim ,
201+ out_features = embedding_dim * 2 ,
202+ use_bias = True ,
203+ dtype = dtype ,
204+ param_dtype = weights_dtype ,
205+ rngs = rngs ,
201206 )
202207
203208 def __call__ (self , x : jax .Array , conditioning_embedding : jax .Array ) -> jax .Array :
@@ -209,7 +214,9 @@ def __call__(self, x: jax.Array, conditioning_embedding: jax.Array) -> jax.Array
209214
210215class NNXAdaLayerNormZero (nnx .Module ):
211216
212- def __init__ (self , embedding_dim : int , eps : float = 1e-6 , dtype : jnp .dtype = jnp .float32 , weights_dtype : jnp .dtype = jnp .float32 ):
217+ def __init__ (
218+ self , embedding_dim : int , eps : float = 1e-6 , dtype : jnp .dtype = jnp .float32 , weights_dtype : jnp .dtype = jnp .float32
219+ ):
213220 self .embedding_dim = embedding_dim
214221 self .eps = eps
215222 self .dtype = dtype
@@ -227,7 +234,9 @@ def __call__(self, x: jax.Array, emb: jax.Array):
227234
228235class NNXAdaLayerNormZeroSingle (nnx .Module ):
229236
230- def __init__ (self , embedding_dim : int , eps : float = 1e-6 , dtype : jnp .dtype = jnp .float32 , weights_dtype : jnp .dtype = jnp .float32 ):
237+ def __init__ (
238+ self , embedding_dim : int , eps : float = 1e-6 , dtype : jnp .dtype = jnp .float32 , weights_dtype : jnp .dtype = jnp .float32
239+ ):
231240 self .embedding_dim = embedding_dim
232241 self .eps = eps
233242 self .dtype = dtype
0 commit comments