Skip to content

Commit efa05ae

Browse files
author
Toshi Pahadia
committed
Optimize WAN Pipeline and VAE Inference
- Implement spatial/temporal parallelism for WAN VAE - Replace resize with jnp.repeat for upsampling - Resolve eager execution regression in VACE pipeline using vae_encode_pass - Add comprehensive sharding validation and fallbacks in VAE with max_logging - Fix multi-host addressable data logic and TeaCache tracking bugs - Optimize pipeline formatting to reduce host overhead - Restore profiler trace dumping functionality in generate_wan.py
1 parent dddc939 commit efa05ae

11 files changed

Lines changed: 231 additions & 130 deletions

‎src/maxdiffusion/generate_wan.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -286,6 +286,8 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None):
286286
# would silently hit stale binaries.
287287
"flash_block_sizes": str(config.flash_block_sizes),
288288
"mesh_shape": str(pipeline.mesh.shape),
289+
"vae_spatial": str(config.vae_spatial),
290+
"vae_decode_chunk": str(config.vae_decode_chunk),
289291
"weights_dtype": str(config.weights_dtype),
290292
"activations_dtype": str(config.activations_dtype),
291293
"scan_layers": str(config.scan_layers),

‎src/maxdiffusion/models/wan/autoencoder_kl_wan.py‎

Lines changed: 58 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323
from jax import tree_util
2424
from flax import nnx
2525
from ...configuration_utils import ConfigMixin
26+
from ... import max_logging
2627
from ..modeling_flax_utils import FlaxModelMixin, get_activation
2728
from ... import common_types
2829
from ..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

205216
class 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

‎src/maxdiffusion/pipelines/wan/wan_pipeline.py‎

Lines changed: 55 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -474,6 +474,38 @@ def __init__(
474474
# repeated serving requests) skip the ~10s/call CPU text encoder.
475475
self._prompt_embeds_cache = {}
476476

477+
def check_inputs(
478+
self,
479+
prompt: Union[str, List[str]] = None,
480+
negative_prompt: Optional[Union[str, List[str]]] = None,
481+
height: int = 480,
482+
width: int = 832,
483+
prompt_embeds: Optional[jax.Array] = None,
484+
negative_prompt_embeds: Optional[jax.Array] = None,
485+
**kwargs,
486+
):
487+
"""Validate user-facing pipeline inputs and shape contracts."""
488+
if prompt is not None and prompt_embeds is not None:
489+
raise ValueError(
490+
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
491+
" only forward one of the two."
492+
)
493+
elif negative_prompt is not None and negative_prompt_embeds is not None:
494+
raise ValueError(
495+
f"Cannot forward both `negative_prompt`: {negative_prompt} and"
496+
f" `negative_prompt_embeds`: {negative_prompt_embeds}. Please make sure to"
497+
" only forward one of the two."
498+
)
499+
500+
mesh = getattr(self, "vae_mesh", getattr(self, "mesh", None))
501+
if mesh is not None and hasattr(mesh, "shape"):
502+
vae_spatial = mesh.shape.get("vae_spatial", 1)
503+
if vae_spatial > 1 and (width // 8) % vae_spatial != 0:
504+
max_logging.log(
505+
f"Warning: Latent width is not divisible by vae_spatial mesh axis ({vae_spatial})."
506+
" VAE spatial sharding will be partially bypassed."
507+
)
508+
477509
@classmethod
478510
def load_text_encoder(cls, config: HyperParameters):
479511
text_encoder_dtype = getattr(config, "text_encoder_dtype", "float32")
@@ -907,8 +939,10 @@ def _decode_latents_to_video(self, latents: jax.Array, trace: Optional[dict] = N
907939
if trace is not None:
908940
trace["vae_decode_tpu"] = time.perf_counter() - t_vae_tpu_start
909941

910-
video = jax.experimental.multihost_utils.process_allgather(video, tiled=True)
911-
video = np.array(video)
942+
if hasattr(video, "addressable_shards") and len(video.addressable_shards) > 0:
943+
video = np.asarray(video.addressable_shards[0].data)
944+
else:
945+
video = np.asarray(video)
912946
return video
913947

914948
@classmethod
@@ -1231,6 +1265,14 @@ def _prepare_model_inputs(
12311265
prompt_embeds: jax.Array = None,
12321266
negative_prompt_embeds: jax.Array = None,
12331267
):
1268+
self.check_inputs(
1269+
prompt=prompt,
1270+
negative_prompt=negative_prompt,
1271+
height=height,
1272+
width=width,
1273+
prompt_embeds=prompt_embeds,
1274+
negative_prompt_embeds=negative_prompt_embeds,
1275+
)
12341276
if max_sequence_length is None:
12351277
max_sequence_length = getattr(self.config, "max_sequence_length", 512)
12361278

@@ -1315,7 +1357,6 @@ def __call__(self, **kwargs):
13151357
aot_cache.cached_jit,
13161358
static_argnames=(
13171359
"do_classifier_free_guidance",
1318-
"guidance_scale",
13191360
"return_residual",
13201361
"skip_blocks",
13211362
),
@@ -1337,6 +1378,8 @@ def transformer_forward_pass(
13371378
rotary_emb=None,
13381379
encoder_attention_mask=None,
13391380
):
1381+
if do_classifier_free_guidance and latents.shape[0] != prompt_embeds.shape[0]:
1382+
latents = jnp.concatenate([latents, latents], axis=0)
13401383
wan_transformer = nnx.merge(graphdef, sharded_state, rest_of_state)
13411384
outputs = wan_transformer(
13421385
hidden_states=latents,
@@ -1362,11 +1405,9 @@ def transformer_forward_pass(
13621405
noise_uncond = noise_pred[bsz:] # Second half = unconditional
13631406
noise_pred = noise_uncond + guidance_scale * (noise_cond - noise_uncond)
13641407

1365-
latents = latents[:bsz]
1366-
13671408
if return_residual:
1368-
return noise_pred, latents, residual_x
1369-
return noise_pred, latents
1409+
return noise_pred, residual_x
1410+
return noise_pred
13701411

13711412

13721413
@aot_cache.cached_jit
@@ -1389,10 +1430,14 @@ def vae_decode_pass(graphdef, state, rest_of_state, latents):
13891430
video = wan_vae.decode(latents, AutoencoderKLWanCache(wan_vae), return_dict=False)[0]
13901431
video = (video / 2.0) + 0.5
13911432
video = jnp.clip(video, 0.0, 1.0)
1392-
return (video * 255.0).astype(jnp.uint8)
1433+
video = (video * 255.0).astype(jnp.uint8)
1434+
if wan_vae.mesh is not None:
1435+
replicated_sharding = NamedSharding(wan_vae.mesh, P())
1436+
video = jax.lax.with_sharding_constraint(video, replicated_sharding)
1437+
return video
13931438

13941439

1395-
@partial(aot_cache.cached_jit, static_argnames=("guidance_scale",))
1440+
@aot_cache.cached_jit
13961441
def transformer_forward_pass_full_cfg(
13971442
graphdef,
13981443
sharded_state,
@@ -1433,7 +1478,7 @@ def transformer_forward_pass_full_cfg(
14331478
return noise_pred_merged, noise_cond, noise_uncond
14341479

14351480

1436-
@partial(aot_cache.cached_jit, static_argnames=("guidance_scale",))
1481+
@aot_cache.cached_jit
14371482
def transformer_forward_pass_cfg_cache(
14381483
graphdef,
14391484
sharded_state,

‎src/maxdiffusion/pipelines/wan/wan_pipeline_2_1.py‎

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -359,23 +359,23 @@ def scan_body(carry, t):
359359
current_latents, current_scheduler_state = carry
360360

361361
if do_cfg:
362-
latents_doubled = jnp.concatenate([current_latents] * 2)
363362
timestep = jnp.broadcast_to(t, bsz * 2)
364-
noise_pred, _, _ = transformer_forward_pass_full_cfg(
363+
noise_pred = transformer_forward_pass(
365364
graphdef,
366365
sharded_state,
367366
rest_of_state,
368-
latents_doubled,
367+
current_latents,
369368
timestep,
370369
prompt_embeds_combined,
370+
do_classifier_free_guidance=True,
371371
guidance_scale=guidance_scale,
372372
kv_cache=kv_cache,
373373
rotary_emb=rotary_emb,
374374
encoder_attention_mask=encoder_attention_mask,
375375
)
376376
else:
377377
timestep = jnp.broadcast_to(t, bsz)
378-
noise_pred, _ = transformer_forward_pass(
378+
noise_pred = transformer_forward_pass(
379379
graphdef,
380380
sharded_state,
381381
rest_of_state,
@@ -422,11 +422,11 @@ def scan_body(carry, t):
422422
skip_warmup,
423423
)
424424

425-
noise_pred, latents, residual_x_cur = transformer_forward_pass(
425+
noise_pred, residual_x_cur = transformer_forward_pass(
426426
graphdef,
427427
sharded_state,
428428
rest_of_state,
429-
jnp.concatenate([latents] * 2) if do_cfg else latents,
429+
latents,
430430
timestep,
431431
prompt_embeds_combined if do_cfg else prompt_cond_embeds,
432432
do_classifier_free_guidance=do_cfg,
@@ -489,7 +489,7 @@ def scan_body(carry, t):
489489

490490
else:
491491
timestep = jnp.broadcast_to(t, bsz)
492-
noise_pred, latents = transformer_forward_pass(
492+
noise_pred = transformer_forward_pass(
493493
graphdef,
494494
sharded_state,
495495
rest_of_state,

0 commit comments

Comments
 (0)