diff --git a/tests/test_dumas_h3_longvideos.py b/tests/test_dumas_h3_longvideos.py index 20f5960..7e46170 100644 --- a/tests/test_dumas_h3_longvideos.py +++ b/tests/test_dumas_h3_longvideos.py @@ -371,6 +371,116 @@ class DumasH3LongVideosHelperTests(unittest.TestCase): self.assertIn("Jon (navy overalls)", shot) self.assertIn("Becca (green coat)", shot) + def test_plan_only_returns_joined_generations_on_script_socket_with_anchor_override(self): + module = self.module + node = module.H3LongVideos() + + class _Clip: + def tokenize(self, text): + return text + + def encode_from_tokens_scheduled(self, tokens): + return tokens + + class _TorchStub: + @staticmethod + def zeros(shape): + return ("zeros", shape) + + original_torch = module.torch + original_vram_gb = module.vram_gb + original_dit_resident_gb = module.dit_resident_gb + original_lora_overhead_gb = module.lora_overhead_gb + original_check_vae_wiring = module.check_vae_wiring + original_check_text_encoder = module.check_text_encoder + original_apply_h3_model_sampling = module.apply_h3_model_sampling + original_sla_pairing = module.sla_pairing + original_lora_hint_notes = module.lora_hint_notes + original_schedule_balance_note = module.schedule_balance_note + original_kernel_backend_note = module.kernel_backend_note + original_audio_scale_note = module.audio_scale_note + original_quant_accel_note = module.quant_accel_note + original_lora_active = module.lora_active + original_resolve_shot_frames = module.resolve_shot_frames + original_plan_beat_frames = module.plan_beat_frames + original_dialogue_fit_warnings = module.dialogue_fit_warnings + original_dialogue_filler_warnings = module.dialogue_filler_warnings + original_distribute_generations = module.distribute_generations + original_continuity_warnings = module.continuity_warnings + original_empty_av_latent = module._empty_av_latent + try: + module.torch = _TorchStub() + module.vram_gb = lambda: (0.0, 0.0) + module.dit_resident_gb = lambda _model: 0.0 + module.lora_overhead_gb = lambda _model: 0.0 + module.check_vae_wiring = lambda *_args, **_kwargs: None + module.check_text_encoder = lambda *_args, **_kwargs: None + module.apply_h3_model_sampling = lambda model, *_args, **_kwargs: (model, "") + module.sla_pairing = lambda *_args, **_kwargs: ("", False, "") + module.lora_hint_notes = lambda *_args, **_kwargs: [] + module.schedule_balance_note = lambda *_args, **_kwargs: "" + module.kernel_backend_note = lambda *_args, **_kwargs: "" + module.audio_scale_note = lambda *_args, **_kwargs: "" + module.quant_accel_note = lambda *_args, **_kwargs: "" + module.lora_active = lambda _model: False + module.resolve_shot_frames = lambda *_args, **_kwargs: (73, "") + module.plan_beat_frames = lambda beats, fps, budget, per_beat=True: ([73] * len(beats), []) + module.dialogue_fit_warnings = lambda *_args, **_kwargs: [] + module.dialogue_filler_warnings = lambda *_args, **_kwargs: [] + module.distribute_generations = lambda anchor, beats, *_args, **_kwargs: [ + f"[Generation 1] {anchor}. {beats[0]}", + f"[Generation 2] {anchor}. {beats[1]}", + ] + module.continuity_warnings = lambda _gens: [] + module._empty_av_latent = lambda *_args, **_kwargs: ({"samples": "latent"},) + + result = node.run( + model=object(), + clip=_Clip(), + vae=object(), + audio_vae=object(), + prompt="Francine stands alone.\n\nFrancine and Frankie walk together.", + resolution="16:9", + steps=20, + cfg=1.0, + sampler_name="res_multistep", + scheduler="simple", + seed=1, + anchor_override="editorial room, soft practical lighting", + character_memory="Francine = white top\nFrankie = black jacket", + plan_only=True, + ) + + self.assertEqual( + result[3], + "[Generation 1] editorial room, soft practical lighting. Francine stands alone.\n---\n" + "[Generation 2] editorial room, soft practical lighting. Francine and Frankie walk together.", + ) + self.assertIn("2 beat(s)", result[2]) + self.assertIn("2 shot(s)", result[2]) + finally: + module.torch = original_torch + module.vram_gb = original_vram_gb + module.dit_resident_gb = original_dit_resident_gb + module.lora_overhead_gb = original_lora_overhead_gb + module.check_vae_wiring = original_check_vae_wiring + module.check_text_encoder = original_check_text_encoder + module.apply_h3_model_sampling = original_apply_h3_model_sampling + module.sla_pairing = original_sla_pairing + module.lora_hint_notes = original_lora_hint_notes + module.schedule_balance_note = original_schedule_balance_note + module.kernel_backend_note = original_kernel_backend_note + module.audio_scale_note = original_audio_scale_note + module.quant_accel_note = original_quant_accel_note + module.lora_active = original_lora_active + module.resolve_shot_frames = original_resolve_shot_frames + module.plan_beat_frames = original_plan_beat_frames + module.dialogue_fit_warnings = original_dialogue_fit_warnings + module.dialogue_filler_warnings = original_dialogue_filler_warnings + module.distribute_generations = original_distribute_generations + module.continuity_warnings = original_continuity_warnings + module._empty_av_latent = original_empty_av_latent + if __name__ == "__main__": unittest.main()