This commit is contained in:
kijai
2024-10-10 15:35:35 +03:00
parent a7033bdd26
commit 8132750127
2 changed files with 12 additions and 28 deletions
+1 -2
View File
@@ -236,12 +236,11 @@ class PyramidFlowTextEncode:
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
model.vae.enable_tiling()
autocastcondition = not model.dtype == torch.float32
autocast_context = torch.autocast(mm.get_autocast_device(device)) if autocastcondition else nullcontext()
model.text_encoder.to(device)
model.text_encoder.to(torch.float16).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)
@@ -42,20 +42,15 @@ 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='bf16', use_gradient_checkpointing=False, return_log=True,
def __init__(self, model_path, model_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_mixed_training=False, use_flash_attn=False,
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,
corrupt_ratio=1/3, interp_condition_pos=True, stages=[1, 2, 4], **kwargs,
):
super().__init__()
if model_dtype == 'bf16':
torch_dtype = torch.bfloat16
elif model_dtype == 'fp16':
torch_dtype = torch.float16
else:
torch_dtype = torch.float32
torch_dtype = model_dtype
self.stages = stages
self.sample_ratios = sample_ratios
@@ -63,24 +58,14 @@ class PyramidDiTForVideoGeneration:
dit_path = os.path.join(model_path, model_variant)
# The dit
if use_mixed_training:
print("using mixed precision training, do not explicitly casting models")
self.dit = PyramidDiffusionMMDiT.from_pretrained(
dit_path, 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,
)
else:
print("using half precision")
self.dit = PyramidDiffusionMMDiT.from_pretrained(
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,
)
self.dit = PyramidDiffusionMMDiT.from_pretrained(
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,
)
# The text encoder
if load_text_encoder: