This commit is contained in:
kijai
2024-10-10 16:00:45 +03:00
parent 8132750127
commit 182d25b88d
2 changed files with 20 additions and 45 deletions
+13 -8
View File
@@ -33,9 +33,9 @@ class DownloadAndLoadPyramidFlowModel:
},
"optional": {
"precision": (["fp16", "fp32", "bf16"],
{"default": "bf16", "tooltip": "official recommendation is that 2b model should be fp16, 5b model should be bf16"}
),
"model_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }),
"text_encoder_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }),
"vae_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }),
#"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,13 +46,16 @@ class DownloadAndLoadPyramidFlowModel:
FUNCTION = "loadmodel"
CATEGORY = "PyramidFlowWrapper"
def loadmodel(self, model, variant, precision):
def loadmodel(self, model, variant, model_dtype, text_encoder_dtype, vae_dtype):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
mm.soft_empty_cache()
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[model_dtype]
text_encoder_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[text_encoder_dtype]
vae_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[vae_dtype]
base_path = folder_paths.get_folder_paths("pyramidflow")[0]
model_path = os.path.join(base_path, model.split("/")[-1])
@@ -70,7 +73,9 @@ class DownloadAndLoadPyramidFlowModel:
model = PyramidDiTForVideoGeneration(
model_path,
dtype,
model_dtype,
text_encoder_dtype,
vae_dtype,
model_variant=variant,
)
@@ -233,14 +238,13 @@ class PyramidFlowTextEncode:
def sample(self, model, positive_prompt, negative_prompt, keep_model_loaded):
mm.soft_empty_cache()
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
autocastcondition = not model.dtype == torch.float32
autocast_context = torch.autocast(mm.get_autocast_device(device)) if autocastcondition else nullcontext()
model.text_encoder.to(torch.float16).to(device)
model.text_encoder.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)
@@ -295,6 +299,7 @@ class PyramidFlowVAEDecode:
self.vae_video_scale_factor = 1 / 3.0986
self.vae.to(device)
latents = latents.to(self.vae.dtype)
if latents.shape[2] == 1:
latents = (latents / self.vae_scale_factor) + self.vae_shift_factor
else:
@@ -42,7 +42,7 @@ 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, use_gradient_checkpointing=False, return_log=True,
def __init__(self, model_path, model_dtype, text_encoder_dtype, vae_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_flash_attn=False,
load_text_encoder=True, load_vae=True, max_temporal_length=31, frame_per_unit=1, use_temporal_causal=True,
@@ -69,13 +69,13 @@ class PyramidDiTForVideoGeneration:
# The text encoder
if load_text_encoder:
self.text_encoder = SD3TextEncoderWithMask(model_path, torch_dtype=torch_dtype)
self.text_encoder = SD3TextEncoderWithMask(model_path, torch_dtype=text_encoder_dtype)
else:
self.text_encoder = None
# The base video vae decoder
if load_vae:
self.vae = CausalVideoVAE.from_pretrained(os.path.join(model_path, 'causal_video_vae'), torch_dtype=torch_dtype, interpolate=False)
self.vae = CausalVideoVAE.from_pretrained(os.path.join(model_path, 'causal_video_vae'), torch_dtype=vae_dtype, interpolate=False)
# Freeze vae
for parameter in self.vae.parameters():
parameter.requires_grad = False
@@ -453,7 +453,6 @@ class PyramidDiTForVideoGeneration:
min_guidance_scale: float = 2.0,
use_linear_guidance: bool = False,
alpha: float = 0.5,
negative_prompt: Optional[Union[str, List[str]]]="cartoon style, worst quality, low quality, blurry, absolute black, absolute white, low res, extra limbs, extra digits, misplaced objects, mutated anatomy, monochrome, horror",
num_images_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
output_type: Optional[str] = "pil",
@@ -511,6 +510,10 @@ class PyramidDiTForVideoGeneration:
pooled_prompt_embeds = torch.cat([negative_pooled_prompt_embeds, positive_pooled_prompt_embeds], dim=0)
prompt_attention_mask = torch.cat([negative_prompt_attention_mask, positive_prompt_attention_mask], dim=0)
prompt_embeds = prompt_embeds.to(dtype)
pooled_prompt_embeds = pooled_prompt_embeds.to(dtype)
prompt_attention_mask = prompt_attention_mask.to(dtype)
# Create the initial random noise
num_channels_latents = self.dit.config.in_channels
latents = self.prepare_latents(
@@ -625,39 +628,6 @@ class PyramidDiTForVideoGeneration:
return image
def decode_latent(self, latents, device):
self.vae.to(device)
if latents.shape[2] == 1:
latents = (latents / self.vae_scale_factor) + self.vae_shift_factor
else:
latents[:, :, :1] = (latents[:, :, :1] / self.vae_scale_factor) + self.vae_shift_factor
latents[:, :, 1:] = (latents[:, :, 1:] / self.vae_video_scale_factor) + self.vae_video_shift_factor
image = self.vae.decode(latents, temporal_chunk=True, window_size=2, tile_sample_min_size=128).sample
self.vae.to('cpu')
image = image.float()
image = (image / 2 + 0.5).clamp(0, 1)
image = rearrange(image, "B C T H W -> (B T) C H W")
image = image.cpu().permute(0, 2, 3, 1).numpy()
image = self.numpy_to_pil(image)
return image
@staticmethod
def numpy_to_pil(images):
"""
Convert a numpy image or a batch of images to a PIL image.
"""
if images.ndim == 3:
images = images[None, ...]
images = (images * 255).round().astype("uint8")
if images.shape[-1] == 1:
# special case for grayscale (single channel) images
pil_images = [Image.fromarray(image.squeeze(), mode="L") for image in images]
else:
pil_images = [Image.fromarray(image) for image in images]
return pil_images
@property
def device(self):
return next(self.dit.parameters()).device