dtypes
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user