diff --git a/nodes.py b/nodes.py index 9f797ff..694bae0 100644 --- a/nodes.py +++ b/nodes.py @@ -33,7 +33,7 @@ class DownloadAndLoadPyramidFlowModel: }, "optional": { - "model_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }), + "model_dtype": (["fp8_e4m3fn","fp8_e5m2","fp16", "fp32", "bf16"],{"default": "bf16", }), "text_encoder_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }), "vae_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }), "use_flash_attn": ("BOOLEAN", {"default": False}), @@ -53,7 +53,7 @@ class DownloadAndLoadPyramidFlowModel: offload_device = mm.unet_offload_device() mm.soft_empty_cache() - model_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[model_dtype] + model_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32, "fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e5m2": torch.float8_e5m2}[model_dtype] text_encoder_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[text_encoder_dtype] vae_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[vae_dtype] @@ -193,12 +193,15 @@ class PyramidFlowSampler: device = mm.get_torch_device() offload_device = mm.unet_offload_device() + dtype = model["dtype"] torch.manual_seed(seed) torch.cuda.manual_seed(seed) - autocastcondition = not model["model"].dtype == torch.float32 - autocast_context = torch.autocast(mm.get_autocast_device(device), dtype=model["model"].dtype) if autocastcondition else nullcontext() + 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() if input_latent is None: with autocast_context: diff --git a/pyramid_dit/modeling_normalization.py b/pyramid_dit/modeling_normalization.py index 6b8361c..dd3da5d 100644 --- a/pyramid_dit/modeling_normalization.py +++ b/pyramid_dit/modeling_normalization.py @@ -57,10 +57,16 @@ class RMSNorm(nn.Module): hidden_states = hidden_states * torch.rsqrt(variance + self.eps) if self.weight is not None: - # convert into half-precision if necessary - if self.weight.dtype in [torch.float16, torch.bfloat16]: - hidden_states = hidden_states.to(self.weight.dtype) - hidden_states = hidden_states * self.weight + # Handle the case where self.weight is torch.float8_e4m3fn + if self.weight.dtype == torch.float8_e4m3fn: + weight = self.weight.to(torch.float32) + else: + weight = self.weight + + # Convert hidden_states to the dtype of weight if necessary + if hidden_states.dtype != weight.dtype: + hidden_states = hidden_states.to(weight.dtype) + hidden_states = hidden_states * weight hidden_states = hidden_states.to(input_dtype) diff --git a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py index db47440..6c9b863 100644 --- a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py +++ b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py @@ -50,7 +50,10 @@ class PyramidDiTForVideoGeneration: ): super().__init__() - torch_dtype = model_dtype + if model_dtype in [torch.float8_e4m3fn, torch.float8_e4m3fn]: + self.dtype = torch.bfloat16 + else: + self.dtype = model_dtype self.stages = stages self.sample_ratios = sample_ratios @@ -61,7 +64,7 @@ class PyramidDiTForVideoGeneration: self.dit = PyramidDiffusionMMDiT.from_pretrained( dit_path, - torch_dtype=torch_dtype, + torch_dtype=self.dtype, use_gradient_checkpointing=use_gradient_checkpointing, use_flash_attn=use_flash_attn, use_t5_mask=True, @@ -70,6 +73,10 @@ class PyramidDiTForVideoGeneration: use_temporal_causal=True if not use_flash_attn else False, interp_condition_pos=interp_condition_pos, ) + if model_dtype in [torch.float8_e4m3fn, torch.float8_e4m3fn]: + for name, param in self.dit.named_parameters(): + if name != "pos_embedding": + param.data = param.data.to(model_dtype) # The text encoder if load_text_encoder: @@ -349,9 +356,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(dtype) + # pooled_prompt_embeds = pooled_prompt_embeds.to(dtype) + # prompt_attention_mask = prompt_attention_mask.to(dtype) # Create the initial random noise @@ -561,7 +568,6 @@ class PyramidDiTForVideoGeneration: generated_latents_list = [] # The generated results last_generated_latents = None - #self.dit.to(torch.float8_e4m3fn) self.dit.to(device) comfy_pbar = ProgressBar(num_units) @@ -649,15 +655,7 @@ class PyramidDiTForVideoGeneration: @property def device(self): return next(self.dit.parameters()).device - - @property - def dtype(self): - return next(self.dit.parameters()).dtype - @property - def vae_dtype(self): - return next(self.dit.parameters()).dtype - @property def guidance_scale(self): return self._guidance_scale