update flash_attn and reqs
This commit is contained in:
@@ -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,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
|
||||||
Reference in New Issue
Block a user