Skip to content

Commit c89f07c

Browse files
committed
Optimize WAN pipeline: hoist timesteps, batched text encoding, I2V concat/transpose churn reduction, replace prints with max_logging, and fix CFG cache multi-host sharding crashes
1 parent 5e884cf commit c89f07c

8 files changed

Lines changed: 124 additions & 74 deletions

File tree

src/maxdiffusion/loaders/ltx2_lora_nnx_loader.py

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -76,8 +76,6 @@ def translate_fn(nnx_path_str):
7676
# the merge_fn warns about unmatched keys in each dict, so we only warn about any leftovers
7777
unmatched_keys = set(h_state_dict) - set(transformer_state_dict) - set(connector_state_dict)
7878
if unmatched_keys:
79-
max_logging.log(
80-
f"{len(unmatched_keys)} key(s) in LoRA dictionary routed to no merge target: {unmatched_keys}"
81-
)
79+
max_logging.log(f"{len(unmatched_keys)} key(s) in LoRA dictionary routed to no merge target: {unmatched_keys}")
8280

8381
return pipeline

src/maxdiffusion/models/quantizations.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@
2525
import jax.numpy as jnp
2626
from jax.tree_util import tree_flatten_with_path, tree_unflatten
2727
from typing import Tuple, Sequence
28+
from maxdiffusion import max_logging
2829

2930
# Params used to define mixed precision quantization configs
3031
DEFAULT = "__default__" # default config
@@ -139,7 +140,7 @@ def _get_quant_config(config):
139140
else:
140141
drhs_bits = 8
141142
drhs_accumulator_dtype = jnp.int32
142-
print(config.quantization_local_shard_count) # -1
143+
max_logging.log(config.quantization_local_shard_count) # -1
143144
drhs_local_aqt = aqt_config.LocalAqt(contraction_axis_shard_count=config.quantization_local_shard_count)
144145
return aqt_config.config_v4(
145146
fwd_bits=8,

src/maxdiffusion/pipelines/wan/wan_pipeline.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -997,7 +997,8 @@ def _prepare_model_inputs_i2v(
997997

998998
prompt_embeds = jax.device_put(prompt_embeds, data_sharding)
999999
negative_prompt_embeds = jax.device_put(negative_prompt_embeds, data_sharding)
1000-
image_embeds = jax.device_put(image_embeds, data_sharding)
1000+
if image_embeds is not None:
1001+
image_embeds = jax.device_put(image_embeds, data_sharding)
10011002

10021003
return prompt_embeds, negative_prompt_embeds, image_embeds, effective_batch_size
10031004

src/maxdiffusion/pipelines/wan/wan_pipeline_2_1.py

Lines changed: 19 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@
2525
import numpy as np
2626
import time
2727
from ... import max_utils
28+
from maxdiffusion import max_logging
2829

2930

3031
class WanPipeline2_1(WanPipeline):
@@ -101,6 +102,7 @@ def __call__(
101102
magcache_K: Optional[int] = None,
102103
retention_ratio: Optional[float] = None,
103104
use_kv_cache: bool = False,
105+
output_type: str = "pil",
104106
):
105107
config = getattr(self, "config", None)
106108
if max_sequence_length is None:
@@ -170,6 +172,9 @@ def __call__(
170172
latents.block_until_ready()
171173
trace["denoise_total"] = time.perf_counter() - t_denoise_start
172174

175+
if output_type == "latent":
176+
return latents, trace
177+
173178
t_decode_start = time.perf_counter()
174179
video = self._decode_latents_to_video(latents, trace=trace)
175180
if hasattr(video, "block_until_ready"):
@@ -222,6 +227,18 @@ def run_inference_2_1(
222227
do_cfg = guidance_scale > 1.0
223228
bsz = latents.shape[0]
224229

230+
data_shards = 1
231+
try:
232+
if hasattr(latents, "sharding") and hasattr(latents.sharding, "mesh"):
233+
data_shards = latents.sharding.mesh.shape["data"] * latents.sharding.mesh.shape.get("fsdp", 1)
234+
except Exception:
235+
pass
236+
237+
if use_cfg_cache and do_cfg and bsz % data_shards != 0:
238+
max_logging.log(
239+
f"Warning: Disabling CFG cache because batch size {bsz} is not divisible by data shards {data_shards}. This often happens with data_parallelism > 1 and per_device_batch_size = 1."
240+
)
241+
use_cfg_cache = False
225242
# Resolution-dependent CFG cache config (FasterCache / MixCache guidance)
226243
if height >= 720:
227244
# 720p: conservative — protect last 40%, interval=5
@@ -306,10 +323,9 @@ def run_inference_2_1(
306323
)
307324

308325
scan_diffusion_loop = getattr(config, "scan_diffusion_loop", False) if config else False
326+
timesteps = jnp.array(scheduler_state.timesteps, dtype=jnp.int32)
309327

310328
if scan_diffusion_loop and not use_magcache and not use_cfg_cache:
311-
timesteps = jnp.array(scheduler_state.timesteps, dtype=jnp.int32)
312-
313329
scheduler_state = scheduler_state.replace(last_sample=jnp.zeros_like(latents), step_index=jnp.array(0, dtype=jnp.int32))
314330

315331
def scan_body(carry, t):
@@ -365,7 +381,7 @@ def scan_body(carry, t):
365381
profiler = max_utils.Profiler(config)
366382
profiler.start()
367383

368-
t = jnp.array(scheduler_state.timesteps, dtype=jnp.int32)[step]
384+
t = timesteps[step]
369385

370386
if use_magcache and do_cfg:
371387
timestep = jnp.broadcast_to(t, bsz * 2 if do_cfg else bsz)

src/maxdiffusion/pipelines/wan/wan_pipeline_2_2.py

Lines changed: 35 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -118,6 +118,7 @@ def __call__(
118118
use_cfg_cache: bool = False,
119119
use_sen_cache: bool = False,
120120
use_kv_cache: bool = False,
121+
output_type: str = "pil",
121122
):
122123
config = getattr(self, "config", None)
123124
if max_sequence_length is None:
@@ -203,6 +204,9 @@ def __call__(
203204
latents.block_until_ready()
204205
trace["denoise_total"] = time.perf_counter() - t_denoise_start
205206

207+
if output_type == "latent":
208+
return latents, trace
209+
206210
t_decode_start = time.perf_counter()
207211
video = self._decode_latents_to_video(latents, trace=trace)
208212
if hasattr(video, "block_until_ready"):
@@ -252,6 +256,19 @@ def run_inference_2_2(
252256
do_classifier_free_guidance = guidance_scale_low > 1.0 or guidance_scale_high > 1.0
253257
bsz = latents.shape[0]
254258

259+
data_shards = 1
260+
try:
261+
if hasattr(latents, "sharding") and hasattr(latents.sharding, "mesh"):
262+
data_shards = latents.sharding.mesh.shape["data"] * latents.sharding.mesh.shape.get("fsdp", 1)
263+
except Exception:
264+
pass
265+
266+
if use_cfg_cache and do_classifier_free_guidance and bsz % data_shards != 0:
267+
max_logging.log(
268+
f"Warning: Disabling CFG cache because batch size {bsz} is not divisible by data shards {data_shards}. This often happens with data_parallelism > 1 and per_device_batch_size = 1."
269+
)
270+
use_cfg_cache = False
271+
255272
prompt_embeds_combined = (
256273
jnp.concatenate([prompt_embeds, negative_prompt_embeds], axis=0) if do_classifier_free_guidance else prompt_embeds
257274
)
@@ -279,6 +296,8 @@ def run_inference_2_2(
279296
high_transformer = nnx.merge(high_noise_graphdef, high_noise_state, high_noise_rest)
280297
kv_cache_high, encoder_attention_mask_high = high_transformer.compute_kv_cache(prompt_embeds_combined)
281298

299+
timesteps = jnp.array(scheduler_state.timesteps, dtype=jnp.int32)
300+
282301
# ── SenCache path (arXiv:2602.24208) ──
283302
if use_sen_cache and do_classifier_free_guidance:
284303
timesteps_np = np.array(scheduler_state.timesteps, dtype=np.int32)
@@ -303,16 +322,18 @@ def run_inference_2_2(
303322
num_train_timesteps = float(scheduler.config.num_train_timesteps)
304323

305324
# SenCache state
306-
ref_noise_pred = None # y^r: cached denoiser output
307-
ref_latent = None # x^r: latent at last cache refresh
308-
ref_timestep = 0.0 # t^r: timestep (normalized to [0,1]) at last cache refresh
309-
accum_dx = 0.0 # accumulated ||Δx|| since last refresh
310-
accum_dt = 0.0 # accumulated |Δt| since last refresh
311-
reuse_count = 0 # consecutive cache reuses
312-
cache_count = 0
325+
ref_noise_pred = jnp.zeros(
326+
(bsz * 2, latents.shape[1], latents.shape[2], latents.shape[3], latents.shape[4]), dtype=latents.dtype
327+
)
328+
ref_latent = jnp.zeros_like(latents)
329+
ref_timestep = jnp.array(0.0, dtype=jnp.float32)
330+
accum_dx = jnp.array(0.0, dtype=jnp.float32)
331+
accum_dt = jnp.array(0.0, dtype=jnp.float32)
332+
reuse_count = jnp.array(0, dtype=jnp.int32)
333+
cache_count = jnp.array(0, dtype=jnp.int32)
313334

314335
for step in range(num_inference_steps):
315-
t = jnp.array(scheduler_state.timesteps, dtype=jnp.int32)[step]
336+
t = timesteps[step]
316337
t_float = float(timesteps_np[step]) / num_train_timesteps # normalize to [0, 1]
317338

318339
# Select transformer and guidance scale
@@ -358,10 +379,10 @@ def run_inference_2_2(
358379
)
359380
ref_noise_pred = noise_pred
360381
ref_latent = latents
361-
ref_timestep = t_float
362-
accum_dx = 0.0
363-
accum_dt = 0.0
364-
reuse_count = 0
382+
ref_timestep = jnp.array(t_float, dtype=jnp.float32)
383+
accum_dx = jnp.array(0.0, dtype=jnp.float32)
384+
accum_dt = jnp.array(0.0, dtype=jnp.float32)
385+
reuse_count = jnp.array(0, dtype=jnp.int32)
365386
latents, scheduler_state = scheduler.step(scheduler_state, noise_pred, t, latents).to_tuple()
366387
continue
367388

@@ -375,12 +396,10 @@ def run_inference_2_2(
375396
score = alpha_x * accum_dx + alpha_t * accum_dt
376397

377398
if score <= sen_epsilon and reuse_count < max_reuse:
378-
# Cache hit: reuse previous output
379399
noise_pred = ref_noise_pred
380400
reuse_count += 1
381401
cache_count += 1
382402
else:
383-
# Cache miss: full CFG forward pass
384403
latents_doubled = jnp.concatenate([latents] * 2)
385404
timestep = jnp.broadcast_to(t, bsz * 2)
386405
noise_pred, _, _ = transformer_forward_pass_full_cfg(
@@ -470,7 +489,7 @@ def run_inference_2_2(
470489
cached_noise_uncond = None
471490

472491
for step in range(num_inference_steps):
473-
t = jnp.array(scheduler_state.timesteps, dtype=jnp.int32)[step]
492+
t = timesteps[step]
474493
is_cache_step = step_is_cache[step]
475494

476495
# Select transformer and guidance scale based on precomputed schedule
@@ -607,8 +626,6 @@ def low_noise_branch(operands):
607626
)
608627

609628
if scan_diffusion_loop:
610-
timesteps = jnp.array(scheduler_state.timesteps, dtype=jnp.int32)
611-
612629
scheduler_state = scheduler_state.replace(last_sample=jnp.zeros_like(latents), step_index=jnp.array(0, dtype=jnp.int32))
613630

614631
def scan_body(carry, t):
@@ -657,7 +674,7 @@ def scan_body(carry, t):
657674
profiler = max_utils.Profiler(config)
658675
profiler.start()
659676

660-
t = jnp.array(scheduler_state.timesteps, dtype=jnp.int32)[step]
677+
t = timesteps[step]
661678

662679
if step_uses_high[step]:
663680
graphdef, state, rest = (

src/maxdiffusion/pipelines/wan/wan_pipeline_i2v_2p1.py

Lines changed: 11 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -213,19 +213,17 @@ def __call__(
213213
last_image,
214214
)
215215

216-
def _process_image_input(img_input, height, width, num_videos_per_prompt):
216+
def _process_image_input(img_input, height, width):
217217
if img_input is None:
218218
return None
219219
tensor = self.video_processor.preprocess(img_input, height=height, width=width)
220220
jax_array = jnp.array(tensor.cpu().numpy())
221221
if jax_array.ndim == 3:
222222
jax_array = jax_array[None, ...] # Add batch dimension
223-
if num_videos_per_prompt > 1:
224-
jax_array = jnp.repeat(jax_array, num_videos_per_prompt, axis=0)
225223
return jax_array
226224

227-
image_tensor = _process_image_input(image, height, width, effective_batch_size)
228-
last_image_tensor = _process_image_input(last_image, height, width, effective_batch_size)
225+
image_tensor = _process_image_input(image, height, width)
226+
last_image_tensor = _process_image_input(last_image, height, width)
229227

230228
if rng is None:
231229
rng = jax.random.key(self.config.seed)
@@ -352,6 +350,8 @@ def run_inference_2_1_i2v(
352350
image_embeds_combined = image_embeds
353351
condition_combined = condition
354352

353+
condition_combined = jnp.transpose(condition_combined, (0, 4, 1, 2, 3))
354+
355355
transformer_obj = nnx.merge(graphdef, sharded_state, rest_of_state)
356356

357357
# Compute RoPE once as it only depends on shape
@@ -373,10 +373,9 @@ def run_inference_2_1_i2v(
373373
)
374374

375375
scan_diffusion_loop = getattr(config, "scan_diffusion_loop", False) if config else False
376+
timesteps = jnp.array(scheduler_state.timesteps, dtype=jnp.int32)
376377

377378
if scan_diffusion_loop and not use_magcache:
378-
timesteps = jnp.array(scheduler_state.timesteps, dtype=jnp.int32)
379-
380379
scheduler_state = scheduler_state.replace(last_sample=jnp.zeros_like(latents), step_index=jnp.array(0, dtype=jnp.int32))
381380

382381
def scan_body(carry, t):
@@ -386,9 +385,9 @@ def scan_body(carry, t):
386385
if do_cfg:
387386
latents_input = jnp.concatenate([current_latents, current_latents], axis=0)
388387

389-
latent_model_input = jnp.concatenate([latents_input, condition_combined], axis=-1)
388+
latents_input = jnp.transpose(latents_input, (0, 4, 1, 2, 3))
389+
latent_model_input = jnp.concatenate([latents_input, condition_combined], axis=1)
390390
timestep = jnp.broadcast_to(t, latents_input.shape[0])
391-
latent_model_input = jnp.transpose(latent_model_input, (0, 4, 1, 2, 3))
392391

393392
outputs = transformer_forward_pass(
394393
graphdef,
@@ -429,7 +428,7 @@ def scan_body(carry, t):
429428
profiler = max_utils.Profiler(config)
430429
profiler.start()
431430

432-
t = jnp.array(scheduler_state.timesteps, dtype=jnp.int32)[step]
431+
t = timesteps[step]
433432

434433
skip_blocks = False
435434
if use_magcache and do_cfg:
@@ -446,9 +445,9 @@ def scan_body(carry, t):
446445
if do_cfg:
447446
latents_input = jnp.concatenate([latents, latents], axis=0)
448447

449-
latent_model_input = jnp.concatenate([latents_input, condition_combined], axis=-1)
448+
latents_input = jnp.transpose(latents_input, (0, 4, 1, 2, 3))
449+
latent_model_input = jnp.concatenate([latents_input, condition_combined], axis=1)
450450
timestep = jnp.broadcast_to(t, latents_input.shape[0])
451-
latent_model_input = jnp.transpose(latent_model_input, (0, 4, 1, 2, 3))
452451

453452
outputs = transformer_forward_pass(
454453
graphdef,

0 commit comments

Comments
 (0)