From 51fb1da7eef0ab69abff2313dc672d0f22d144ce Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Fri, 11 Oct 2024 00:30:55 +0300 Subject: [PATCH] dtype fixes --- nodes.py | 3 +-- pyramid_dit/pyramid_dit_for_video_gen_pipeline.py | 15 ++++++--------- 2 files changed, 7 insertions(+), 11 deletions(-) diff --git a/nodes.py b/nodes.py index 694bae0..86debb2 100644 --- a/nodes.py +++ b/nodes.py @@ -199,7 +199,6 @@ class PyramidFlowSampler: torch.cuda.manual_seed(seed) autocast_dtype = dtype if dtype not in [torch.float8_e4m3fn, torch.float8_e5m2] else torch.bfloat16 - print(autocast_dtype) autocastcondition = not dtype == torch.float32 autocast_context = torch.autocast(mm.get_autocast_device(device), dtype=autocast_dtype) if autocastcondition else nullcontext() @@ -266,7 +265,7 @@ class PyramidFlowTextEncode: text_encoder = model["model"].text_encoder autocastcondition = not model["text_encoder_dtype"] == torch.float32 - autocast_context = torch.autocast(mm.get_autocast_device(device)) if autocastcondition else nullcontext() + autocast_context = torch.autocast(mm.get_autocast_device(device), dtype=model["text_encoder_dtype"]) if autocastcondition else nullcontext() text_encoder.to(device) with autocast_context: diff --git a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py index 6c9b863..f12fe78 100644 --- a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py +++ b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py @@ -483,9 +483,6 @@ class PyramidDiTForVideoGeneration: output_type: Optional[str] = "pil", device: Optional[torch.device] = None, ): - #device = self.device - dtype = self.dtype - assert (temp - 1) % self.frame_per_unit == 0, "The frames should be divided by frame_per unit" # if isinstance(prompt, str): @@ -535,9 +532,9 @@ class PyramidDiTForVideoGeneration: pooled_prompt_embeds = torch.cat([negative_pooled_prompt_embeds, positive_pooled_prompt_embeds], dim=0) prompt_attention_mask = torch.cat([negative_prompt_attention_mask, positive_prompt_attention_mask], dim=0) - # prompt_embeds = prompt_embeds.to(dtype) - # pooled_prompt_embeds = pooled_prompt_embeds.to(dtype) - # prompt_attention_mask = prompt_attention_mask.to(dtype) + prompt_embeds = prompt_embeds.to(self.dtype) + pooled_prompt_embeds = pooled_prompt_embeds.to(self.dtype) + prompt_attention_mask = prompt_attention_mask.to(self.dtype) # Create the initial random noise num_channels_latents = self.dit.config.in_channels @@ -547,7 +544,7 @@ class PyramidDiTForVideoGeneration: temp, height, width, - prompt_embeds.dtype, + self.dtype, device, generator, ) @@ -589,7 +586,7 @@ class PyramidDiTForVideoGeneration: width, 1, device, - dtype, + self.dtype, generator, is_first_frame=True, ) @@ -634,7 +631,7 @@ class PyramidDiTForVideoGeneration: width, self.frame_per_unit, device, - dtype, + self.dtype, generator, is_first_frame=False, )