diff --git a/nodes.py b/nodes.py index eb82d5f..95c45d6 100644 --- a/nodes.py +++ b/nodes.py @@ -33,9 +33,9 @@ class DownloadAndLoadPyramidFlowModel: }, "optional": { - "precision": (["fp16", "fp32", "bf16"], - {"default": "bf16", "tooltip": "official recommendation is that 2b model should be fp16, 5b model should be bf16"} - ), + "model_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }), + "text_encoder_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }), + "vae_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }), #"fp8_transformer": (['disabled', 'enabled', 'fastmode'], {"default": 'disabled', "tooltip": "enabled casts the transformer to torch.float8_e4m3fn, fastmode is only for latest nvidia GPUs"}), #"compile": (["disabled","onediff","torch"], {"tooltip": "compile the model for faster inference, these are advanced options only available on Linux, see readme for more info"}), } @@ -46,13 +46,16 @@ class DownloadAndLoadPyramidFlowModel: FUNCTION = "loadmodel" CATEGORY = "PyramidFlowWrapper" - def loadmodel(self, model, variant, precision): + def loadmodel(self, model, variant, model_dtype, text_encoder_dtype, vae_dtype): device = mm.get_torch_device() offload_device = mm.unet_offload_device() mm.soft_empty_cache() - dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + model_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[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] + base_path = folder_paths.get_folder_paths("pyramidflow")[0] model_path = os.path.join(base_path, model.split("/")[-1]) @@ -70,7 +73,9 @@ class DownloadAndLoadPyramidFlowModel: model = PyramidDiTForVideoGeneration( model_path, - dtype, + model_dtype, + text_encoder_dtype, + vae_dtype, model_variant=variant, ) @@ -233,14 +238,13 @@ class PyramidFlowTextEncode: def sample(self, model, positive_prompt, negative_prompt, keep_model_loaded): mm.soft_empty_cache() - device = mm.get_torch_device() offload_device = mm.unet_offload_device() autocastcondition = not model.dtype == torch.float32 autocast_context = torch.autocast(mm.get_autocast_device(device)) if autocastcondition else nullcontext() - model.text_encoder.to(torch.float16).to(device) + model.text_encoder.to(device) with autocast_context: prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = model.text_encoder(positive_prompt, device) negative_prompt_embeds, negative_prompt_attention_mask, pooled_negative_prompt_embeds = model.text_encoder(negative_prompt, device) @@ -295,6 +299,7 @@ class PyramidFlowVAEDecode: self.vae_video_scale_factor = 1 / 3.0986 self.vae.to(device) + latents = latents.to(self.vae.dtype) if latents.shape[2] == 1: latents = (latents / self.vae_scale_factor) + self.vae_shift_factor else: diff --git a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py index 34fd730..e930ddc 100644 --- a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py +++ b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py @@ -42,7 +42,7 @@ class PyramidDiTForVideoGeneration: The pyramid dit for both image and video generation, The running class wrapper This class is mainly for fixed unit implementation: 1 + n + n + n """ - def __init__(self, model_path, model_dtype, use_gradient_checkpointing=False, return_log=True, + def __init__(self, model_path, model_dtype, text_encoder_dtype, vae_dtype, use_gradient_checkpointing=False, return_log=True, model_variant="diffusion_transformer_768p", timestep_shift=1.0, stage_range=[0, 1/3, 2/3, 1], sample_ratios=[1, 1, 1], scheduler_gamma=1/3, use_flash_attn=False, load_text_encoder=True, load_vae=True, max_temporal_length=31, frame_per_unit=1, use_temporal_causal=True, @@ -69,13 +69,13 @@ class PyramidDiTForVideoGeneration: # The text encoder if load_text_encoder: - self.text_encoder = SD3TextEncoderWithMask(model_path, torch_dtype=torch_dtype) + self.text_encoder = SD3TextEncoderWithMask(model_path, torch_dtype=text_encoder_dtype) else: self.text_encoder = None # The base video vae decoder if load_vae: - self.vae = CausalVideoVAE.from_pretrained(os.path.join(model_path, 'causal_video_vae'), torch_dtype=torch_dtype, interpolate=False) + self.vae = CausalVideoVAE.from_pretrained(os.path.join(model_path, 'causal_video_vae'), torch_dtype=vae_dtype, interpolate=False) # Freeze vae for parameter in self.vae.parameters(): parameter.requires_grad = False @@ -453,7 +453,6 @@ class PyramidDiTForVideoGeneration: min_guidance_scale: float = 2.0, use_linear_guidance: bool = False, alpha: float = 0.5, - negative_prompt: Optional[Union[str, List[str]]]="cartoon style, worst quality, low quality, blurry, absolute black, absolute white, low res, extra limbs, extra digits, misplaced objects, mutated anatomy, monochrome, horror", num_images_per_prompt: Optional[int] = 1, generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, output_type: Optional[str] = "pil", @@ -511,6 +510,10 @@ 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) + # Create the initial random noise num_channels_latents = self.dit.config.in_channels latents = self.prepare_latents( @@ -625,39 +628,6 @@ class PyramidDiTForVideoGeneration: return image - def decode_latent(self, latents, device): - self.vae.to(device) - if latents.shape[2] == 1: - latents = (latents / self.vae_scale_factor) + self.vae_shift_factor - else: - latents[:, :, :1] = (latents[:, :, :1] / self.vae_scale_factor) + self.vae_shift_factor - latents[:, :, 1:] = (latents[:, :, 1:] / self.vae_video_scale_factor) + self.vae_video_shift_factor - - image = self.vae.decode(latents, temporal_chunk=True, window_size=2, tile_sample_min_size=128).sample - self.vae.to('cpu') - image = image.float() - image = (image / 2 + 0.5).clamp(0, 1) - image = rearrange(image, "B C T H W -> (B T) C H W") - image = image.cpu().permute(0, 2, 3, 1).numpy() - image = self.numpy_to_pil(image) - return image - - @staticmethod - def numpy_to_pil(images): - """ - Convert a numpy image or a batch of images to a PIL image. - """ - if images.ndim == 3: - images = images[None, ...] - images = (images * 255).round().astype("uint8") - if images.shape[-1] == 1: - # special case for grayscale (single channel) images - pil_images = [Image.fromarray(image.squeeze(), mode="L") for image in images] - else: - pil_images = [Image.fromarray(image) for image in images] - - return pil_images - @property def device(self): return next(self.dit.parameters()).device