dtype fixes

This commit is contained in:
Jukka Seppänen
2024-10-11 00:30:55 +03:00
parent 88f4d14666
commit 51fb1da7ee
2 changed files with 7 additions and 11 deletions
+1 -2
View File
@@ -199,7 +199,6 @@ class PyramidFlowSampler:
torch.cuda.manual_seed(seed) torch.cuda.manual_seed(seed)
autocast_dtype = dtype if dtype not in [torch.float8_e4m3fn, torch.float8_e5m2] else torch.bfloat16 autocast_dtype = dtype if dtype not in [torch.float8_e4m3fn, torch.float8_e5m2] else torch.bfloat16
print(autocast_dtype)
autocastcondition = not dtype == torch.float32 autocastcondition = not dtype == torch.float32
autocast_context = torch.autocast(mm.get_autocast_device(device), dtype=autocast_dtype) if autocastcondition else nullcontext() 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 text_encoder = model["model"].text_encoder
autocastcondition = not model["text_encoder_dtype"] == torch.float32 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) text_encoder.to(device)
with autocast_context: with autocast_context:
@@ -483,9 +483,6 @@ class PyramidDiTForVideoGeneration:
output_type: Optional[str] = "pil", output_type: Optional[str] = "pil",
device: Optional[torch.device] = None, 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" assert (temp - 1) % self.frame_per_unit == 0, "The frames should be divided by frame_per unit"
# if isinstance(prompt, str): # 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) 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_attention_mask = torch.cat([negative_prompt_attention_mask, positive_prompt_attention_mask], dim=0)
# prompt_embeds = prompt_embeds.to(dtype) prompt_embeds = prompt_embeds.to(self.dtype)
# pooled_prompt_embeds = pooled_prompt_embeds.to(dtype) pooled_prompt_embeds = pooled_prompt_embeds.to(self.dtype)
# prompt_attention_mask = prompt_attention_mask.to(dtype) prompt_attention_mask = prompt_attention_mask.to(self.dtype)
# Create the initial random noise # Create the initial random noise
num_channels_latents = self.dit.config.in_channels num_channels_latents = self.dit.config.in_channels
@@ -547,7 +544,7 @@ class PyramidDiTForVideoGeneration:
temp, temp,
height, height,
width, width,
prompt_embeds.dtype, self.dtype,
device, device,
generator, generator,
) )
@@ -589,7 +586,7 @@ class PyramidDiTForVideoGeneration:
width, width,
1, 1,
device, device,
dtype, self.dtype,
generator, generator,
is_first_frame=True, is_first_frame=True,
) )
@@ -634,7 +631,7 @@ class PyramidDiTForVideoGeneration:
width, width,
self.frame_per_unit, self.frame_per_unit,
device, device,
dtype, self.dtype,
generator, generator,
is_first_frame=False, is_first_frame=False,
) )