dtype fixes
This commit is contained in:
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user