update flash_attn and reqs

This commit is contained in:
Jukka Seppänen
2024-10-10 22:10:45 +03:00
parent fa672acdff
commit 754ef838d2
3 changed files with 10 additions and 8 deletions
+2 -2
View File
@@ -36,6 +36,7 @@ class DownloadAndLoadPyramidFlowModel:
"model_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }), "model_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }),
"text_encoder_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }), "text_encoder_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }),
"vae_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"}), #"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"}), #"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" FUNCTION = "loadmodel"
CATEGORY = "PyramidFlowWrapper" 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() device = mm.get_torch_device()
offload_device = mm.unet_offload_device() offload_device = mm.unet_offload_device()
@@ -196,7 +197,6 @@ class PyramidFlowSampler:
torch.cuda.manual_seed(seed) torch.cuda.manual_seed(seed)
autocastcondition = not model["model"].dtype == torch.float32 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() autocast_context = torch.autocast(mm.get_autocast_device(device), dtype=model["model"].dtype) if autocastcondition else nullcontext()
if input_latent is None: if input_latent is None:
@@ -60,11 +60,15 @@ class PyramidDiTForVideoGeneration:
self.dit = PyramidDiffusionMMDiT.from_pretrained( self.dit = PyramidDiffusionMMDiT.from_pretrained(
dit_path, torch_dtype=torch_dtype, dit_path,
torch_dtype=torch_dtype,
use_gradient_checkpointing=use_gradient_checkpointing, use_gradient_checkpointing=use_gradient_checkpointing,
use_flash_attn=use_flash_attn, use_t5_mask=True, use_flash_attn=use_flash_attn,
add_temp_pos_embed=True, temp_pos_embed_type='rope', use_t5_mask=True,
use_temporal_causal=use_temporal_causal, interp_condition_pos=interp_condition_pos, 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 # The text encoder
-2
View File
@@ -2,7 +2,5 @@ diffusers>=0.30.1
accelerate>=0.30.0 accelerate>=0.30.0
einops einops
packaging packaging
pandas
scikit-image
sentencepiece sentencepiece
timm>=0.6.12 timm>=0.6.12