2323from jax import tree_util
2424from flax import nnx
2525from ...configuration_utils import ConfigMixin
26+ from ... import max_logging
2627from ..modeling_flax_utils import FlaxModelMixin , get_activation
2728from ... import common_types
2829from ..vae_flax import (
@@ -99,9 +100,9 @@ def __init__(
99100 self .mesh = mesh
100101
101102 # Weight sharding (Kernel is sharded along output channels)
102- num_fsdp_devices = mesh .shape ["vae_spatial" ]
103+ num_fsdp_devices = mesh .shape ["vae_spatial" ] if mesh is not None and "vae_spatial" in mesh . shape else 1
103104 kernel_sharding = (None , None , None , None , None )
104- if out_channels % num_fsdp_devices == 0 :
105+ if num_fsdp_devices > 1 and out_channels % num_fsdp_devices == 0 :
105106 kernel_sharding = (None , None , None , None , "vae_spatial" )
106107
107108 self .conv = nnx .Conv (
@@ -119,8 +120,10 @@ def __init__(
119120 )
120121
121122 def __call__ (self , x : jax .Array , cache_x : Optional [jax .Array ] = None , idx = - 1 ) -> jax .Array :
122- spatial_sharding = NamedSharding (self .mesh , P ("redundant" , None , None , "vae_spatial" , None ))
123- x = jax .lax .with_sharding_constraint (x , spatial_sharding )
123+ if self .mesh is not None and "vae_spatial" in self .mesh .shape :
124+ spatial_sharding = NamedSharding (self .mesh , P ("redundant" , None , None , "vae_spatial" , None ))
125+ if spatial_sharding is not None :
126+ x = jax .lax .with_sharding_constraint (x , spatial_sharding )
124127
125128 current_padding = list (self ._causal_padding )
126129 padding_needed = self ._depth_padding_before
@@ -198,8 +201,16 @@ def __call__(self, x: jax.Array) -> jax.Array:
198201 n , h , w , c = in_shape
199202 target_h = int (h * self .scale_factor [0 ])
200203 target_w = int (w * self .scale_factor [1 ])
201- out = jax .image .resize (x .astype (jnp .float32 ), (n , target_h , target_w , c ), method = self .method )
202- return out .astype (input_dtype )
204+ if self .method == "nearest" and self .scale_factor [0 ] == int (self .scale_factor [0 ]) and self .scale_factor [1 ] == int (self .scale_factor [1 ]):
205+ scale_h = int (self .scale_factor [0 ])
206+ scale_w = int (self .scale_factor [1 ])
207+ out = jnp .repeat (jnp .repeat (x , scale_h , axis = 1 ), scale_w , axis = 2 )
208+ else :
209+ if self .method == "nearest" :
210+ max_logging .log (f"Warning: WanUpsample2D nearest method requested but scale_factor { self .scale_factor } is not integer. Falling back to jax.image.resize." )
211+ out = jax .image .resize (x .astype (jnp .float32 ), (n , target_h , target_w , c ), method = self .method )
212+ out = out .astype (input_dtype )
213+ return out
203214
204215
205216class Identity (nnx .Module ):
@@ -225,14 +236,16 @@ def __init__(
225236 weights_dtype : jnp .dtype = jnp .float32 ,
226237 precision : jax .lax .Precision = None ,
227238 ):
239+ rank = len (kernel_size ) if isinstance (kernel_size , (tuple , list )) else 2
240+ kernel_sharding = (None ,) * (rank + 2 )
228241 self .conv = nnx .Conv (
229242 dim ,
230243 dim ,
231244 kernel_size = kernel_size ,
232245 strides = stride ,
233246 use_bias = True ,
234247 rngs = rngs ,
235- kernel_init = nnx .with_partitioning (nnx .initializers .xavier_uniform (), ( None , None , None , None ) ),
248+ kernel_init = nnx .with_partitioning (nnx .initializers .xavier_uniform (), kernel_sharding ),
236249 dtype = dtype ,
237250 param_dtype = weights_dtype ,
238251 precision = precision ,
@@ -1131,7 +1144,6 @@ def __init__(
11311144 )
11321145 self .mesh = mesh
11331146
1134- @nnx .jit
11351147 def _encode (self , x : jax .Array , feat_cache : AutoencoderKLWanCache ):
11361148 feat_cache .init_cache ()
11371149 if x .shape [- 1 ] != 3 :
@@ -1151,7 +1163,11 @@ def _encode(self, x: jax.Array, feat_cache: AutoencoderKLWanCache):
11511163 iter_ = 1 + ((t - 1 + CHUNK_SIZE - 1 ) // CHUNK_SIZE ) if t > 1 else 1
11521164 enc_feat_map = feat_cache ._enc_feat_map
11531165
1154- spatial_sharding = NamedSharding (self .mesh , P ("redundant" , None , None , "vae_spatial" , None ))
1166+ spatial_sharding = (
1167+ NamedSharding (self .mesh , P ("redundant" , None , None , "vae_spatial" , None ))
1168+ if self .mesh is not None and "vae_spatial" in self .mesh .shape
1169+ else None
1170+ )
11551171
11561172 def finalize (out , enc_feat_map ):
11571173 feat_cache ._enc_feat_map = enc_feat_map
@@ -1162,7 +1178,8 @@ def finalize(out, enc_feat_map):
11621178 with jax .named_scope ("AutoencoderKLWan_encode_chunk_0" ):
11631179 chunk_0 = x [:, :1 , ...]
11641180 out_0 , enc_feat_map , _ = self .encoder (chunk_0 , feat_cache = enc_feat_map , feat_idx = 0 )
1165- out_0 = jax .lax .with_sharding_constraint (out_0 , spatial_sharding )
1181+ if spatial_sharding is not None :
1182+ out_0 = jax .lax .with_sharding_constraint (out_0 , spatial_sharding )
11661183
11671184 if iter_ <= 1 :
11681185 return finalize (out_0 , enc_feat_map )
@@ -1172,11 +1189,13 @@ def finalize(out, enc_feat_map):
11721189 with jax .named_scope ("AutoencoderKLWan_encode_chunk_1" ):
11731190 chunk_1 = x [:, 1 : (1 + CHUNK_SIZE ), ...]
11741191 out_1 , enc_feat_map , _ = self .encoder (chunk_1 , feat_cache = enc_feat_map , feat_idx = 0 )
1175- out_1 = jax .lax .with_sharding_constraint (out_1 , spatial_sharding )
1192+ if spatial_sharding is not None :
1193+ out_1 = jax .lax .with_sharding_constraint (out_1 , spatial_sharding )
11761194
11771195 if iter_ <= 2 :
11781196 out = jnp .concatenate ([out_0 , out_1 ], axis = 1 )
1179- out = jax .lax .with_sharding_constraint (out , spatial_sharding )
1197+ if spatial_sharding is not None :
1198+ out = jax .lax .with_sharding_constraint (out , spatial_sharding )
11801199 return finalize (out , enc_feat_map )
11811200
11821201 # Prepare the remaining chunks to be scanned over
@@ -1209,10 +1228,13 @@ def finalize(out, enc_feat_map):
12091228 def scan_fn (carry , chunk ):
12101229 current_feat_map = carry
12111230 local_encoder = nnx .merge (graphdef , state )
1231+ if spatial_sharding is not None :
1232+ chunk = jax .lax .with_sharding_constraint (chunk , spatial_sharding )
12121233 out_chunk , next_feat_map , _ = local_encoder (chunk , feat_cache = current_feat_map , feat_idx = 0 )
1213- out_chunk = jax .lax .with_sharding_constraint (out_chunk , spatial_sharding )
1234+ if spatial_sharding is not None :
1235+ out_chunk = jax .lax .with_sharding_constraint (out_chunk , spatial_sharding )
12141236 next_feat_map = jax .tree_util .tree_map (
1215- lambda x : jax .lax .with_sharding_constraint (x , spatial_sharding ) if isinstance (x , jax . Array ) else x , next_feat_map
1237+ lambda x : jax .lax .with_sharding_constraint (x , spatial_sharding ) if spatial_sharding is not None and hasattr (x , "shape" ) and x . ndim == len ( spatial_sharding . spec ) else x , next_feat_map
12161238 )
12171239 return next_feat_map , out_chunk
12181240
@@ -1225,7 +1247,8 @@ def scan_fn(carry, chunk):
12251247 out_rest = out_rest [:, : T_rest // self .temporal_downsample_factor , ...]
12261248
12271249 out = jnp .concatenate ([out_0 , out_1 , out_rest ], axis = 1 )
1228- out = jax .lax .with_sharding_constraint (out , spatial_sharding )
1250+ if spatial_sharding is not None :
1251+ out = jax .lax .with_sharding_constraint (out , spatial_sharding )
12291252 return finalize (out , enc_feat_map )
12301253
12311254 @jax .named_scope ("AutoencoderKLWan_encode" )
@@ -1239,7 +1262,6 @@ def encode(
12391262 return (posterior ,)
12401263 return FlaxAutoencoderKLOutput (latent_dist = posterior )
12411264
1242- @nnx .jit
12431265 def _decode (
12441266 self , z : jax .Array , feat_cache : AutoencoderKLWanCache , return_dict : bool = True
12451267 ) -> Union [FlaxDecoderOutput , jax .Array ]:
@@ -1249,20 +1271,28 @@ def _decode(
12491271 x = self .post_quant_conv (z )
12501272
12511273 dec_feat_map = feat_cache ._feat_map
1252- spatial_sharding = NamedSharding (self .mesh , P ("redundant" , None , None , "vae_spatial" , None ))
1274+ spatial_sharding = (
1275+ NamedSharding (self .mesh , P ("redundant" , None , None , "vae_spatial" , None ))
1276+ if self .mesh is not None and "vae_spatial" in self .mesh .shape
1277+ else None
1278+ )
12531279
12541280 # First chunk (i=0)
12551281 with jax .named_scope ("AutoencoderKLWan_decode_chunk_0" ):
1256- chunk_in_0 = jax .lax .with_sharding_constraint (x [:, 0 :1 , ...], spatial_sharding )
1282+ if spatial_sharding is not None :
1283+ chunk_in_0 = jax .lax .with_sharding_constraint (x [:, 0 :1 , ...], spatial_sharding )
12571284 out_0 , dec_feat_map , _ = self .decoder (chunk_in_0 , feat_cache = dec_feat_map , feat_idx = 0 )
1258- out_0 = jax .lax .with_sharding_constraint (out_0 , spatial_sharding )
1285+ if spatial_sharding is not None :
1286+ out_0 = jax .lax .with_sharding_constraint (out_0 , spatial_sharding )
12591287
12601288 if iter_ > 1 :
12611289 # Run chunk 1 outside scan to properly form the cache shape
12621290 with jax .named_scope ("AutoencoderKLWan_decode_chunk_1" ):
1263- chunk_in_1 = jax .lax .with_sharding_constraint (x [:, 1 :2 , ...], spatial_sharding )
1291+ if spatial_sharding is not None :
1292+ chunk_in_1 = jax .lax .with_sharding_constraint (x [:, 1 :2 , ...], spatial_sharding )
12641293 out_chunk_1 , dec_feat_map , _ = self .decoder (chunk_in_1 , feat_cache = dec_feat_map , feat_idx = 0 )
1265- out_chunk_1 = jax .lax .with_sharding_constraint (out_chunk_1 , spatial_sharding )
1294+ if spatial_sharding is not None :
1295+ out_chunk_1 = jax .lax .with_sharding_constraint (out_chunk_1 , spatial_sharding )
12661296
12671297 out_1 = out_chunk_1
12681298 out_list = [out_0 , out_1 ]
@@ -1297,11 +1327,13 @@ def _decode(
12971327 def scan_fn (carry , chunk_in ):
12981328 current_feat_map = carry
12991329 local_decoder = nnx .merge (graphdef , state )
1300- chunk_in = jax .lax .with_sharding_constraint (chunk_in , spatial_sharding )
1330+ if spatial_sharding is not None :
1331+ chunk_in = jax .lax .with_sharding_constraint (chunk_in , spatial_sharding )
13011332 out_chunk , next_feat_map , _ = local_decoder (chunk_in , feat_cache = current_feat_map , feat_idx = 0 )
1302- out_chunk = jax .lax .with_sharding_constraint (out_chunk , spatial_sharding )
1333+ if spatial_sharding is not None :
1334+ out_chunk = jax .lax .with_sharding_constraint (out_chunk , spatial_sharding )
13031335 next_feat_map = jax .tree_util .tree_map (
1304- lambda x : jax .lax .with_sharding_constraint (x , spatial_sharding ) if isinstance (x , jax . Array ) else x ,
1336+ lambda x : jax .lax .with_sharding_constraint (x , spatial_sharding ) if spatial_sharding is not None and hasattr (x , "shape" ) and x . ndim == len ( spatial_sharding . spec ) else x ,
13051337 next_feat_map ,
13061338 )
13071339 return next_feat_map , out_chunk
@@ -1314,7 +1346,8 @@ def scan_fn(carry, chunk_in):
13141346 out_list .append (out_rest )
13151347
13161348 out = jnp .concatenate (out_list , axis = 1 )
1317- out = jax .lax .with_sharding_constraint (out , spatial_sharding )
1349+ if spatial_sharding is not None :
1350+ out = jax .lax .with_sharding_constraint (out , spatial_sharding )
13181351 else :
13191352 out = out_0
13201353
0 commit comments