linoyts HF Staff commited on
Commit
4f01ac3
·
verified ·
1 Parent(s): 3af1c47

Add mask_drop_step for single-pass harmonize

Browse files
Files changed (1) hide show
  1. minimax_h3_inpaint_blocks.py +30 -11
minimax_h3_inpaint_blocks.py CHANGED
@@ -655,6 +655,11 @@ class MiniMaxH3InpaintSetTimestepsStep(MiniMaxH3SetTimestepsStep):
655
  InputParam(
656
  name="inpaint_audio_row_mask", type_hint=torch.Tensor, required=True, description="Per audio row."
657
  ),
 
 
 
 
 
658
  ]
659
 
660
  @staticmethod
@@ -701,6 +706,8 @@ class MiniMaxH3InpaintSetTimestepsStep(MiniMaxH3SetTimestepsStep):
701
  block_state.timesteps = components.scheduler.timesteps
702
  block_state.audio_timesteps = components.audio_scheduler.timesteps
703
 
 
 
704
  block_state.row_timestep_plan = [
705
  tuple(
706
  tensor.to(device)
@@ -714,11 +721,12 @@ class MiniMaxH3InpaintSetTimestepsStep(MiniMaxH3SetTimestepsStep):
714
  float(audio_timestep),
715
  max(float(timestep), components.keyframe_noise_aug),
716
  AUDIO_COND_TIMESTEP,
717
- block_state.inpaint_row_mask,
 
718
  block_state.inpaint_audio_row_mask,
719
  )
720
  )
721
- for timestep, audio_timestep in zip(block_state.timesteps, block_state.audio_timesteps)
722
  ]
723
 
724
  self.set_block_state(state, block_state)
@@ -756,6 +764,7 @@ class MiniMaxH3InpaintLoopSchedulerStep(MiniMaxH3LoopSchedulerStep):
756
  InputParam(
757
  name="inpaint_audio_row_mask", type_hint=torch.Tensor, required=True, description="Per audio row."
758
  ),
 
759
  ]
760
 
761
  @torch.no_grad()
@@ -765,17 +774,27 @@ class MiniMaxH3InpaintLoopSchedulerStep(MiniMaxH3LoopSchedulerStep):
765
  # On the last step the schedule has reached the clean end, and so does the source: the preserved rows are
766
  # handed to the decoder as the footage itself rather than as the anchor's 0.999.
767
  last = i + 1 >= block_state.timesteps.numel()
768
- video_level = 1.0 if last else VISUAL_COND_TIMESTEP
769
 
770
  num_condition_video_rows = block_state.num_condition_video_rows
771
- block_state.latents[num_condition_video_rows:] = impose_source(
772
- block_state.latents[num_condition_video_rows:],
773
- block_state.source_rows,
774
- block_state.inpaint_noise_rows,
775
- block_state.inpaint_row_mask,
776
- video_level,
777
- components.scheduler,
778
- )
 
 
 
 
 
 
 
 
 
 
779
  if block_state.source_audio_rows is not None:
780
  num_condition_audio_rows = block_state.num_condition_audio_rows
781
  block_state.audio_latents[num_condition_audio_rows:] = impose_source(
 
655
  InputParam(
656
  name="inpaint_audio_row_mask", type_hint=torch.Tensor, required=True, description="Per audio row."
657
  ),
658
+ InputParam(
659
+ name="mask_drop_step", type_hint=int,
660
+ description="If set, the video mask is dropped (all rows become targets) from this step on, so the "
661
+ "preserved subject free-refines into the scene over the final steps. None keeps the mask all the way.",
662
+ ),
663
  ]
664
 
665
  @staticmethod
 
706
  block_state.timesteps = components.scheduler.timesteps
707
  block_state.audio_timesteps = components.audio_scheduler.timesteps
708
 
709
+ drop = block_state.mask_drop_step
710
+ ones_video = torch.ones_like(block_state.inpaint_row_mask)
711
  block_state.row_timestep_plan = [
712
  tuple(
713
  tensor.to(device)
 
721
  float(audio_timestep),
722
  max(float(timestep), components.keyframe_noise_aug),
723
  AUDIO_COND_TIMESTEP,
724
+ # after the drop step every video row is a free target; audio stays preserved either way
725
+ (ones_video if (drop is not None and i >= drop) else block_state.inpaint_row_mask),
726
  block_state.inpaint_audio_row_mask,
727
  )
728
  )
729
+ for i, (timestep, audio_timestep) in enumerate(zip(block_state.timesteps, block_state.audio_timesteps))
730
  ]
731
 
732
  self.set_block_state(state, block_state)
 
764
  InputParam(
765
  name="inpaint_audio_row_mask", type_hint=torch.Tensor, required=True, description="Per audio row."
766
  ),
767
+ InputParam(name="mask_drop_step", type_hint=int, description="See the set-timesteps step."),
768
  ]
769
 
770
  @torch.no_grad()
 
774
  # On the last step the schedule has reached the clean end, and so does the source: the preserved rows are
775
  # handed to the decoder as the footage itself rather than as the anchor's 0.999.
776
  last = i + 1 >= block_state.timesteps.numel()
777
+ drop = block_state.mask_drop_step
778
 
779
  num_condition_video_rows = block_state.num_condition_video_rows
780
+
781
+ # Video write-back. With a mask-drop, stop imposing the source from step `drop` on (the subject then
782
+ # free-refines into the scene), and on the handoff step `drop - 1` re-noise the preserved rows to step
783
+ # `drop`'s level so the next forward reads them on distribution. Without a drop, the normal inpaint write-back.
784
+ if drop is None or i < drop:
785
+ if drop is not None and i == drop - 1:
786
+ video_level = float(block_state.timesteps[drop])
787
+ else:
788
+ video_level = 1.0 if last else VISUAL_COND_TIMESTEP
789
+ block_state.latents[num_condition_video_rows:] = impose_source(
790
+ block_state.latents[num_condition_video_rows:],
791
+ block_state.source_rows,
792
+ block_state.inpaint_noise_rows,
793
+ block_state.inpaint_row_mask,
794
+ video_level,
795
+ components.scheduler,
796
+ )
797
+ # Audio is always preserved: we harmonize the picture, not the soundtrack.
798
  if block_state.source_audio_rows is not None:
799
  num_condition_audio_rows = block_state.num_condition_audio_rows
800
  block_state.audio_latents[num_condition_audio_rows:] = impose_source(