diff --git a/nodes.py b/nodes.py index ac6d9c2..5bb1834 100644 --- a/nodes.py +++ b/nodes.py @@ -36,6 +36,7 @@ class DownloadAndLoadPyramidFlowModel: "model_dtype": (["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}), #"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,7 +47,7 @@ class DownloadAndLoadPyramidFlowModel: FUNCTION = "loadmodel" CATEGORY = "PyramidFlowWrapper" - def loadmodel(self, model, variant, model_dtype, text_encoder_dtype, vae_dtype): + def loadmodel(self, model, variant, model_dtype, text_encoder_dtype, vae_dtype, use_flash_attn=False): device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -196,7 +197,6 @@ class PyramidFlowSampler: torch.cuda.manual_seed(seed) autocastcondition = not model["model"].dtype == torch.float32 - #autocastcondition = True autocast_context = torch.autocast(mm.get_autocast_device(device), dtype=model["model"].dtype) if autocastcondition else nullcontext() if input_latent is None: diff --git a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py index 0dfad62..db47440 100644 --- a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py +++ b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py @@ -60,11 +60,15 @@ class PyramidDiTForVideoGeneration: self.dit = PyramidDiffusionMMDiT.from_pretrained( - dit_path, torch_dtype=torch_dtype, + dit_path, + torch_dtype=torch_dtype, use_gradient_checkpointing=use_gradient_checkpointing, - use_flash_attn=use_flash_attn, use_t5_mask=True, - add_temp_pos_embed=True, temp_pos_embed_type='rope', - use_temporal_causal=use_temporal_causal, interp_condition_pos=interp_condition_pos, + use_flash_attn=use_flash_attn, + use_t5_mask=True, + add_temp_pos_embed=True, + temp_pos_embed_type='rope', + use_temporal_causal=True if not use_flash_attn else False, + interp_condition_pos=interp_condition_pos, ) # The text encoder diff --git a/requirements.txt b/requirements.txt index 6f571e5..19c6c4c 100644 --- a/requirements.txt +++ b/requirements.txt @@ -2,7 +2,5 @@ diffusers>=0.30.1 accelerate>=0.30.0 einops packaging -pandas -scikit-image sentencepiece timm>=0.6.12 \ No newline at end of file