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", }),
"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:
@@ -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
-2
View File
@@ -2,7 +2,5 @@ diffusers>=0.30.1
accelerate>=0.30.0
einops
packaging
pandas
scikit-image
sentencepiece
timm>=0.6.12