[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:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user