[fix]: magi-human SR latent prep: invalidate stale packed layout

C4 (4190c720) added a precompute_static_packed_layout call in the base
latent prep stage that stashes coords/modality-maps/max_ch on
batch.magi_static_packed_layout, sized for the BASE-resolution latent.

The SR latent prep stage upsamples batch.latents to the SR grid (e.g.
256x480 -> 512x896 for SR-540p), changing video_token_num and the
shapes of video_coords / video_mm — but it didn't invalidate the
precomputed layout. The SR denoising loop then passed the stale
base-sized layout to build_static_packed_inputs, which produced a
modality_mapping whose first-dim mismatched the SR-sized token tensor,
crashing in MagiHumanDiT.adapter at the text_mask scatter:

  IndexError: The shape of the mask [3243] at index 0 does not match
  the shape of the indexed tensor [11771, 3584] at index 0

Fix: clear batch.magi_static_packed_layout in MagiHumanSRLatentPrep so
the SR denoising loop falls back to the slow path of
build_static_packed_inputs (which rebuilds from the current latent
shape). Base C4 perf win is preserved (32 base steps); SR has only ~5
steps so the meshgrid recompute cost is negligible.

Repro: examples/inference/basic/basic_magi_human_sr540p.py now runs
end-to-end (34s on B200). SR-540p parity tests t2v + ti2v + DiT parity
+ distill DiT parity all pass.
This commit is contained in:
will
2026-05-05 15:47:43 -07:00
parent 990d2c2410
commit f1eeb6303f
@@ -156,6 +156,16 @@ class MagiHumanSRLatentPreparationStage(PipelineStage):
batch.magi_latent_T = latent_t
batch.magi_latent_H = latent_h
batch.magi_latent_W = latent_w
# Invalidate the static packed layout precomputed by the base
# latent prep stage: SR upsamples `batch.latents` to a larger
# spatial grid, which changes video_token_num / video_coords /
# video_mm. The SR denoising loop's
# `getattr(batch, "magi_static_packed_layout", None)` will then
# fall back to the slow path of `build_static_packed_inputs`,
# which rebuilds those fields from the new latent shape. SR
# only does ~5 denoising steps so the meshgrid recompute cost
# is negligible relative to SR-DiT forward.
batch.magi_static_packed_layout = None
if getattr(batch, "image_latent", None) is not None:
batch.image_latent = self._encode_image(batch, actual_h, actual_w)