Big update, refactor model loading, breaks backwards compatibility
archiving old version in legacy branch, probably not going to support SD3 version going forward
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
import os
|
||||
import torch
|
||||
import folder_paths
|
||||
import json
|
||||
import comfy.model_management as mm
|
||||
from comfy.utils import ProgressBar, load_torch_file
|
||||
|
||||
@@ -16,30 +17,116 @@ script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
if not "pyramidflow" in folder_paths.folder_names_and_paths:
|
||||
folder_paths.add_model_folder_path("pyramidflow", os.path.join(folder_paths.models_dir, "pyramidflow"))
|
||||
|
||||
class DownloadAndLoadPyramidFlowModel:
|
||||
|
||||
from .pyramid_dit.mmdit_modules import PyramidDiffusionMMDiT
|
||||
from .pyramid_dit.flux_modules import PyramidFluxTransformer
|
||||
from .video_vae.modeling_causal_vae import CausalVideoVAE
|
||||
|
||||
from contextlib import nullcontext
|
||||
try:
|
||||
from accelerate import init_empty_weights
|
||||
from accelerate.utils import set_module_tensor_to_device
|
||||
is_accelerate_available = True
|
||||
except:
|
||||
is_accelerate_available = False
|
||||
|
||||
class PyramidFlowTorchCompileSettings:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"backend": (["inductor","cudagraphs"], {"default": "inductor"}),
|
||||
"fullgraph": ("BOOLEAN", {"default": False, "tooltip": "Enable full graph mode"}),
|
||||
"mode": (["default", "max-autotune", "max-autotune-no-cudagraphs", "reduce-overhead"], {"default": "default"}),
|
||||
"compile_whole_model": ("BOOLEAN", {"default": False, "tooltip": "Compile the whole model, overrides other block settings"}),
|
||||
"single_blocks": ("BOOLEAN", {"default": True, "tooltip": "Compile single_blocks"}),
|
||||
"double_blocks": ("BOOLEAN", {"default": True, "tooltip": "Compile transformer blocks"}),
|
||||
"embedders": ("BOOLEAN", {"default": True, "tooltip": "Compile embedders"}),
|
||||
"compile_rest": ("BOOLEAN", {"default": True, "tooltip": "Compile the rest of the model (proj and norm out)"}),
|
||||
},
|
||||
}
|
||||
RETURN_TYPES = ("MOCHICOMPILEARGS",)
|
||||
RETURN_NAMES = ("torch_compile_args",)
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "MochiWrapper"
|
||||
DESCRIPTION = "torch.compile settings, when connected to the model loader, torch.compile of the selected layers is attempted. Requires Triton and torch 2.5.0 is recommended"
|
||||
|
||||
def loadmodel(self, backend, fullgraph, mode, compile_whole_model, single_blocks, double_blocks, embedders, compile_rest):
|
||||
|
||||
compile_args = {
|
||||
"backend": backend,
|
||||
"fullgraph": fullgraph,
|
||||
"mode": mode,
|
||||
"compile_whole_model": compile_whole_model,
|
||||
"single_blocks": single_blocks,
|
||||
"double_blocks": double_blocks,
|
||||
"embedders": embedders,
|
||||
"compile_rest": compile_rest,
|
||||
}
|
||||
|
||||
return (compile_args, )
|
||||
|
||||
class PyramidFlowVAELoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": (
|
||||
[
|
||||
"rain1011/pyramid-flow-sd3",
|
||||
"rain1011/pyramid-flow-miniflux"
|
||||
|
||||
],
|
||||
),
|
||||
"variant": (
|
||||
["diffusion_transformer_384p", "diffusion_transformer_768p"],
|
||||
),
|
||||
|
||||
"vae": (folder_paths.get_filename_list("vae"), {"tooltip": "The name of the checkpoint (model) to load.",}),
|
||||
"precision": (["fp16", "bf16", "fp32"], {"default": "bf16"}),
|
||||
},
|
||||
"optional": {
|
||||
"model_dtype": (["fp8_e4m3fn","fp8_e5m2","fp16", "fp32", "bf16"],{"default": "bf16", }),
|
||||
"text_encoder_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }),
|
||||
"vae_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }),
|
||||
"fp8_fastmode": ("BOOLEAN",{"default": False, "tooltip": "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_args": ("MOCHICOMPILEARGS", {"tooltip": "Optional torch.compile arguments",}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("PYRAMIDFLOWVAE", )
|
||||
RETURN_NAMES = ("pyramidflow_vae",)
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "PyramidFlowWrapper"
|
||||
|
||||
def loadmodel(self, vae, precision, compile_args=None):
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
vae_path = folder_paths.get_full_path_or_raise("vae", vae)
|
||||
|
||||
dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||
|
||||
config_path = os.path.join(script_directory, 'configs', 'causal_video_vae_config.json')
|
||||
with open(config_path) as f:
|
||||
config = json.load(f)
|
||||
|
||||
with (init_empty_weights() if is_accelerate_available else nullcontext()):
|
||||
vae = CausalVideoVAE.from_config(config, torch_dtype=dtype, interpolate=False)
|
||||
vae_sd = load_torch_file(vae_path)
|
||||
if is_accelerate_available:
|
||||
for name, param in vae.named_parameters():
|
||||
set_module_tensor_to_device(vae, name, dtype=dtype, device=device, value=vae_sd[name])
|
||||
else:
|
||||
vae.load_state_dict(vae_sd)
|
||||
del vae_sd
|
||||
# Freeze vae
|
||||
for parameter in vae.parameters():
|
||||
parameter.requires_grad = False
|
||||
|
||||
vae.eval().to(device)
|
||||
#torch.compile
|
||||
if compile_args is not None:
|
||||
vae = torch.compile(vae, fullgraph=compile_args["fullgraph"], dynamic=False, backend=compile_args["backend"])
|
||||
|
||||
return (vae,)
|
||||
|
||||
class PyramidFlowModelLoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "The name of the checkpoint (model) to load.",}),
|
||||
"precision": (["fp8_e4m3fn","fp8_e4m3fn_fast","fp16", "fp32", "bf16"], {"default": "bf16"}),
|
||||
},
|
||||
"optional": {
|
||||
"compile_args": ("MOCHICOMPILEARGS", {"tooltip": "Optional torch.compile arguments",}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -48,111 +135,90 @@ class DownloadAndLoadPyramidFlowModel:
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "PyramidFlowWrapper"
|
||||
|
||||
def loadmodel(self, model, variant, model_dtype, text_encoder_dtype, vae_dtype, fp8_fastmode):
|
||||
def loadmodel(self, model, precision, compile_args=None):
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
mm.soft_empty_cache()
|
||||
|
||||
model_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32, "fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e5m2": torch.float8_e5m2}[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]
|
||||
model_path = folder_paths.get_full_path_or_raise("diffusion_models", model)
|
||||
|
||||
base_path = folder_paths.get_folder_paths("pyramidflow")[0]
|
||||
dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||
|
||||
transformer_sd = load_torch_file(model_path)
|
||||
|
||||
for key in transformer_sd:
|
||||
if key.startswith("pos_embed."):
|
||||
model_name = "pyramid_mmdit"
|
||||
continue
|
||||
else:
|
||||
model_name = "pyramid_flux"
|
||||
|
||||
model_path = os.path.join(base_path, model.split("/")[-1])
|
||||
variant_path = os.path.join(model_path, variant)
|
||||
if model_name == "pyramid_flux":
|
||||
config_path = os.path.join(script_directory, 'configs', 'miniflux_transformer_config.json')
|
||||
with open(config_path) as f:
|
||||
config = json.load(f)
|
||||
|
||||
with (init_empty_weights() if is_accelerate_available else nullcontext()):
|
||||
transformer = PyramidFluxTransformer.from_config(config)
|
||||
|
||||
if is_accelerate_available:
|
||||
logging.info("Using accelerate to load and assign model weights to device...")
|
||||
for name, param in transformer.named_parameters():
|
||||
set_module_tensor_to_device(transformer, name, dtype=dtype, device=device, value=transformer_sd[name])
|
||||
else:
|
||||
transformer.load_state_dict(transformer_sd)
|
||||
transformer = transformer.to(dtype)
|
||||
|
||||
elif model_name == "pyramid_mmdit":
|
||||
config_path = os.path.join(script_directory, 'configs', 'mmdit_transformer_config.json')
|
||||
with open(config_path) as f:
|
||||
config = json.load(f)
|
||||
transformer = PyramidDiffusionMMDiT.from_config(config)
|
||||
params_to_keep = {"pos_embedding"}
|
||||
if is_accelerate_available:
|
||||
logging.info("Using accelerate to load and assign model weights to device...")
|
||||
for name, param in transformer.named_parameters():
|
||||
if not any(keyword in name for keyword in params_to_keep):
|
||||
set_module_tensor_to_device(transformer, name, dtype=dtype, device=device, value=transformer_sd[name])
|
||||
else:
|
||||
set_module_tensor_to_device(transformer, name, dtype=torch.bfloat16, device=device, value=transformer_sd[name])
|
||||
else:
|
||||
transformer.load_state_dict(transformer_sd)
|
||||
if dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
|
||||
for name, param in transformer.named_parameters():
|
||||
if not any(keyword in name for keyword in params_to_keep):
|
||||
param.data = param.data.to(dtype)
|
||||
|
||||
if precision == "fp8_e4m3fn_fast":
|
||||
from .fp8_optimization import convert_fp8_linear
|
||||
convert_fp8_linear(transformer, torch.bfloat16)
|
||||
|
||||
if not os.path.exists(variant_path):
|
||||
from huggingface_hub import snapshot_download
|
||||
log.info(f"Downloading model to: {model_path}")
|
||||
ignore_patterns = []
|
||||
if model == "rain1011/pyramid-flow-miniflux":
|
||||
ignore_patterns.extend["*text_encoder*", "*tokenizer*"]
|
||||
if variant == "diffusion_transformer_384p":
|
||||
snapshot_download(
|
||||
repo_id=model,
|
||||
ignore_patterns=["*diffusion_transformer_768p*"],
|
||||
local_dir=model_path,
|
||||
local_dir_use_symlinks=False,
|
||||
)
|
||||
elif variant == "diffusion_transformer_768p":
|
||||
snapshot_download(
|
||||
repo_id=model,
|
||||
ignore_patterns=["*diffusion_transformer_384p*"],
|
||||
local_dir=model_path,
|
||||
local_dir_use_symlinks=False,
|
||||
)
|
||||
model_name = "pyramid_flux" if "flux" in model else "pyramid_mmdit"
|
||||
print(model_name)
|
||||
model = PyramidDiTForVideoGeneration(
|
||||
model_path,
|
||||
model_dtype,
|
||||
model_name,
|
||||
text_encoder_dtype,
|
||||
vae_dtype,
|
||||
model_variant=variant,
|
||||
fp8_fastmode=fp8_fastmode,
|
||||
)
|
||||
transformer.to(device)
|
||||
#torch.compile
|
||||
if compile_args is not None:
|
||||
torch._dynamo.config.force_parameter_static_shapes = False
|
||||
dynamic = True # because of the stages the compiliation should be dynamic
|
||||
if compile_args["compile_whole_model"]:
|
||||
transformer = torch.compile(transformer, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
|
||||
else:
|
||||
if compile_args["single_blocks"]:
|
||||
for i, block in enumerate(transformer.single_transformer_blocks):
|
||||
transformer.single_transformer_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
|
||||
if compile_args["double_blocks"]:
|
||||
for i, block in enumerate(transformer.transformer_blocks):
|
||||
transformer.transformer_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
|
||||
if compile_args["embedders"]:
|
||||
transformer.context_embedder = torch.compile(transformer.context_embedder, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
|
||||
transformer.time_text_embed = torch.compile(transformer.time_text_embed, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
|
||||
transformer.x_embedder = torch.compile(transformer.x_embedder, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
|
||||
if compile_args["compile_rest"]:
|
||||
transformer.norm_out.linear = torch.compile(transformer.norm_out.linear, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
|
||||
transformer.proj_out = torch.compile(transformer.proj_out, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
|
||||
|
||||
|
||||
# # compilation
|
||||
# if compile == "torch":
|
||||
# torch._dynamo.config.suppress_errors = True
|
||||
# pipe.transformer.to(memory_format=torch.channels_last)
|
||||
# pipe.transformer = torch.compile(pipe.transformer, mode="max-autotune", fullgraph=True)
|
||||
# elif compile == "onediff":
|
||||
# from onediffx import compile_pipe
|
||||
# os.environ['NEXFORT_FX_FORCE_TRITON_SDPA'] = '1'
|
||||
|
||||
# pipe = compile_pipe(
|
||||
# pipe,
|
||||
# backend="nexfort",
|
||||
# options= {"mode": "max-optimize:max-autotune:max-autotune", "memory_format": "channels_last", "options": {"inductor.optimize_linear_epilogue": False, "triton.fuse_attention_allow_fp16_reduction": False}},
|
||||
# ignores=["vae"],
|
||||
# fuse_qkv_projections=True if pab_config is None else False,
|
||||
# )
|
||||
pyramid_model = PyramidDiTForVideoGeneration(transformer, dtype, model_name)
|
||||
|
||||
pyramid_pipe = {
|
||||
"model": model,
|
||||
"dtype": model_dtype,
|
||||
"text_encoder_dtype": text_encoder_dtype,
|
||||
"vae_dtype": vae_dtype,
|
||||
}
|
||||
return (pyramid_pipe,)
|
||||
|
||||
|
||||
class CogVideoTextEncode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"clip": ("CLIP",),
|
||||
"prompt": ("STRING", {"default": "", "multiline": True} ),
|
||||
},
|
||||
"optional": {
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"force_offload": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
RETURN_NAMES = ("conditioning",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "CogVideoWrapper"
|
||||
|
||||
def process(self, clip, prompt, strength=1.0, force_offload=True):
|
||||
load_device = mm.text_encoder_device()
|
||||
offload_device = mm.text_encoder_offload_device()
|
||||
clip.tokenizer.t5xxl.pad_to_max_length = True
|
||||
clip.tokenizer.t5xxl.max_length = 226
|
||||
clip.cond_stage_model.to(load_device)
|
||||
tokens = clip.tokenize(prompt, return_word_ids=True)
|
||||
|
||||
embeds = clip.encode_from_tokens(tokens, return_pooled=False, return_dict=False)
|
||||
embeds *= strength
|
||||
if force_offload:
|
||||
clip.cond_stage_model.to(offload_device)
|
||||
|
||||
return (embeds, )
|
||||
return (pyramid_model,)
|
||||
|
||||
class PyramidFlowSampler:
|
||||
@classmethod
|
||||
@@ -177,8 +243,8 @@ class PyramidFlowSampler:
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("PYRAMIDFLOWMODEL", "LATENT", )
|
||||
RETURN_NAMES = ("model","samples", )
|
||||
RETURN_TYPES = ("LATENT", )
|
||||
RETURN_NAMES = ("samples", )
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "PyramidFlowWrapper"
|
||||
|
||||
@@ -188,7 +254,14 @@ class PyramidFlowSampler:
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
dtype = model["dtype"]
|
||||
|
||||
if isinstance(model, dict):
|
||||
pyramid_model = model["model"]
|
||||
#dtype = model["dtype"]
|
||||
else:
|
||||
pyramid_model = model
|
||||
|
||||
dtype = pyramid_model.dit.dtype
|
||||
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
@@ -202,7 +275,7 @@ class PyramidFlowSampler:
|
||||
|
||||
if input_latent is None:
|
||||
with autocast_context:
|
||||
latents = model["model"].generate(
|
||||
latents = pyramid_model.generate(
|
||||
prompt_embeds_dict = prompt_embeds,
|
||||
device=device,
|
||||
num_inference_steps=first_frame_steps,
|
||||
@@ -216,11 +289,11 @@ class PyramidFlowSampler:
|
||||
)
|
||||
else:
|
||||
with autocast_context:
|
||||
latents = model["model"].generate_i2v(
|
||||
latents = pyramid_model.generate_i2v(
|
||||
prompt_embeds_dict = prompt_embeds,
|
||||
input_image_latent=input_latent,
|
||||
device=device,
|
||||
num_inference_steps=video_steps, #why's this a list
|
||||
num_inference_steps=video_steps,
|
||||
height=height,
|
||||
width=width,
|
||||
temp=temp,
|
||||
@@ -228,78 +301,19 @@ class PyramidFlowSampler:
|
||||
output_type="latent",
|
||||
)
|
||||
|
||||
|
||||
if not keep_model_loaded:
|
||||
model["model"].dit.to(offload_device)
|
||||
pyramid_model.dit.to(offload_device)
|
||||
|
||||
return (model, {"samples": latents},)
|
||||
|
||||
return ({"samples": latents},)
|
||||
|
||||
#todo: sd3 version
|
||||
class PyramidFlowTextEncode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("PYRAMIDFLOWMODEL",),
|
||||
"positive_prompt": ("STRING", {"default": "hyper quality, Ultra HD, 8K", "multiline": True} ),
|
||||
"negative_prompt": ("STRING", {"default": "", "multiline": True} ),
|
||||
"keep_model_loaded": ("BOOLEAN", {"default": False}),
|
||||
|
||||
},
|
||||
"optional": {
|
||||
"prev_prompt": ("PYRAMIDFLOWPROMPT", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("PYRAMIDFLOWPROMPT", )
|
||||
RETURN_NAMES = ("prompt_embeds", )
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "PyramidFlowWrapper"
|
||||
|
||||
def sample(self, model, positive_prompt, negative_prompt, keep_model_loaded, prev_prompt=None):
|
||||
mm.soft_empty_cache()
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
text_encoder = model["model"].text_encoder
|
||||
|
||||
autocastcondition = not model["text_encoder_dtype"] == torch.float32
|
||||
autocast_context = torch.autocast(mm.get_autocast_device(device), dtype=model["text_encoder_dtype"]) if autocastcondition else nullcontext()
|
||||
|
||||
text_encoder.to(device)
|
||||
with autocast_context:
|
||||
prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = text_encoder(positive_prompt, device)
|
||||
negative_prompt_embeds, negative_prompt_attention_mask, pooled_negative_prompt_embeds = text_encoder(negative_prompt, device)
|
||||
if not keep_model_loaded:
|
||||
text_encoder.to(offload_device)
|
||||
|
||||
if prev_prompt is not None:
|
||||
prompt_embeds = torch.cat((prev_prompt["prompt_embeds"], prompt_embeds), dim=0)
|
||||
prompt_attention_mask = torch.cat((prev_prompt["attention_mask"], prompt_attention_mask), dim=0)
|
||||
pooled_prompt_embeds = torch.cat((prev_prompt["pooled_embeds"], pooled_prompt_embeds), dim=0)
|
||||
|
||||
negative_prompt_embeds = torch.cat((prev_prompt["negative_prompt_embeds"], negative_prompt_embeds), dim=0)
|
||||
negative_prompt_attention_mask = torch.cat((prev_prompt["negative_attention_mask"], negative_prompt_attention_mask), dim=0)
|
||||
pooled_negative_prompt_embeds = torch.cat((prev_prompt["negative_pooled_embeds"], pooled_negative_prompt_embeds), dim=0)
|
||||
|
||||
embeds = {
|
||||
"prompt_embeds": prompt_embeds,
|
||||
"attention_mask": prompt_attention_mask,
|
||||
"pooled_embeds": pooled_prompt_embeds,
|
||||
"negative_prompt_embeds": negative_prompt_embeds,
|
||||
"negative_attention_mask": negative_prompt_attention_mask,
|
||||
"negative_pooled_embeds": pooled_negative_prompt_embeds
|
||||
}
|
||||
|
||||
return (embeds,)
|
||||
|
||||
#not functional yet, todo: figure out why the results are bad with it
|
||||
class PyramidFlowTextEncodeComfy:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"clip": ("CLIP",),
|
||||
"positive_prompt": ("STRING", {"default": "hyper quality, Ultra HD, 8K", "multiline": True} ),
|
||||
"negative_prompt": ("STRING", {"default": "", "multiline": True} ),
|
||||
"negative_prompt": ("STRING", {"default": "cartoon style, worst quality, low quality, blurry, absolute black, absolute white, low res, extra limbs, extra digits, misplaced objects, mutated anatomy, monochrome, horror", "multiline": True} ),
|
||||
"force_offload": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
}
|
||||
@@ -324,7 +338,6 @@ class PyramidFlowTextEncodeComfy:
|
||||
clip.cond_stage_model.t5_attention_mask = True
|
||||
|
||||
clip.cond_stage_model.to(device)#.to(torch.bfloat16)
|
||||
clip.cond_stage_model.clip_l.to(device)
|
||||
|
||||
#positive
|
||||
tokens = clip.tokenizer.t5xxl.tokenize_with_weights(positive_prompt, return_word_ids=False)
|
||||
@@ -355,8 +368,9 @@ class PyramidFlowVAEEncode:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("PYRAMIDFLOWMODEL",),
|
||||
"image": ("IMAGE",),
|
||||
"vae": ("PYRAMIDFLOWVAE",),
|
||||
"image": ("IMAGE",),
|
||||
"enable_tiling": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -365,19 +379,20 @@ class PyramidFlowVAEEncode:
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "PyramidFlowWrapper"
|
||||
|
||||
def sample(self, model, image):
|
||||
def sample(self, vae, image, enable_tiling):
|
||||
mm.soft_empty_cache()
|
||||
|
||||
self.vae = model["model"].vae
|
||||
dtype = model["vae_dtype"]
|
||||
|
||||
dtype = vae.dtype
|
||||
if enable_tiling:
|
||||
vae.enable_tiling()
|
||||
else:
|
||||
vae.disable_tiling()
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
self.vae.disable_tiling()
|
||||
|
||||
# For the image latent
|
||||
self.vae_shift_factor = 0.1490
|
||||
self.vae_scale_factor = 1 / 1.8415
|
||||
vae_shift_factor = 0.1490
|
||||
vae_scale_factor = 1 / 1.8415
|
||||
|
||||
normalize = transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))
|
||||
input_image_tensor = rearrange(image, 'b h w c -> b c h w')
|
||||
@@ -385,11 +400,10 @@ class PyramidFlowVAEEncode:
|
||||
input_image_tensor = input_image_tensor.unsqueeze(2) # Add temporal dimension t=1
|
||||
input_image_tensor = input_image_tensor.to(dtype=dtype, device=device)
|
||||
|
||||
self.vae.to(device)
|
||||
input_image_latent = (self.vae.encode(input_image_tensor).latent_dist.sample() - self.vae_shift_factor) * self.vae_scale_factor # [b c 1 h w]
|
||||
self.vae.to(offload_device)
|
||||
vae.to(device)
|
||||
input_image_latent = (vae.encode(input_image_tensor).latent_dist.sample() - vae_shift_factor) * vae_scale_factor # [b c 1 h w]
|
||||
vae.to(offload_device)
|
||||
|
||||
|
||||
return (input_image_latent,)
|
||||
|
||||
class PyramidFlowVAEDecode:
|
||||
@@ -397,11 +411,11 @@ class PyramidFlowVAEDecode:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("PYRAMIDFLOWMODEL",),
|
||||
"vae": ("PYRAMIDFLOWVAE",),
|
||||
"samples": ("LATENT",),
|
||||
"tile_sample_min_size": ("INT", {"default": 256, "min": 64, "max": 512, "step": 8}),
|
||||
"window_size": ("INT", {"default": 2, "min": 1, "max": 4, "step": 1}),
|
||||
|
||||
"enable_tiling": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -410,35 +424,37 @@ class PyramidFlowVAEDecode:
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "PyramidFlowWrapper"
|
||||
|
||||
def sample(self, model, samples, tile_sample_min_size, window_size):
|
||||
def sample(self, vae, samples, tile_sample_min_size, window_size, enable_tiling):
|
||||
mm.soft_empty_cache()
|
||||
|
||||
latents = samples["samples"]
|
||||
self.vae = model["model"].vae
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
self.vae.enable_tiling()
|
||||
if enable_tiling:
|
||||
vae.enable_tiling()
|
||||
else:
|
||||
vae.disable_tiling()
|
||||
|
||||
# For the image latent
|
||||
self.vae_shift_factor = 0.1490
|
||||
self.vae_scale_factor = 1 / 1.8415
|
||||
vae_shift_factor = 0.1490
|
||||
vae_scale_factor = 1 / 1.8415
|
||||
|
||||
# For the video latent
|
||||
self.vae_video_shift_factor = -0.2343
|
||||
self.vae_video_scale_factor = 1 / 3.0986
|
||||
vae_video_shift_factor = -0.2343
|
||||
vae_video_scale_factor = 1 / 3.0986
|
||||
|
||||
self.vae.to(device)
|
||||
latents = latents.to(self.vae.dtype)
|
||||
vae.to(device)
|
||||
latents = latents.to(vae.dtype)
|
||||
if latents.shape[2] == 1:
|
||||
latents = (latents / self.vae_scale_factor) + self.vae_shift_factor
|
||||
latents = (latents / vae_scale_factor) + 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
|
||||
latents[:, :, :1] = (latents[:, :, :1] / vae_scale_factor) + vae_shift_factor
|
||||
latents[:, :, 1:] = (latents[:, :, 1:] / vae_video_scale_factor) + vae_video_shift_factor
|
||||
|
||||
image = self.vae.decode(latents, temporal_chunk=True, window_size=window_size, tile_sample_min_size=tile_sample_min_size).sample
|
||||
image = vae.decode(latents, temporal_chunk=True, window_size=window_size, tile_sample_min_size=tile_sample_min_size).sample
|
||||
|
||||
self.vae.to(offload_device)
|
||||
vae.to(offload_device)
|
||||
|
||||
image = image.float()
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
@@ -450,12 +466,13 @@ class PyramidFlowVAEDecode:
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"DownloadAndLoadPyramidFlowModel": DownloadAndLoadPyramidFlowModel,
|
||||
"PyramidFlowSampler": PyramidFlowSampler,
|
||||
"PyramidFlowVAEDecode": PyramidFlowVAEDecode,
|
||||
"PyramidFlowTextEncode": PyramidFlowTextEncode,
|
||||
"PyramidFlowVAEEncode": PyramidFlowVAEEncode,
|
||||
"PyramidFlowTextEncodeComfy": PyramidFlowTextEncodeComfy,
|
||||
"PyramidFlowTorchCompileSettings": PyramidFlowTorchCompileSettings,
|
||||
"PyramidFlowTransformerLoader": PyramidFlowModelLoader,
|
||||
"PyramidFlowVAELoader": PyramidFlowVAELoader
|
||||
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -464,5 +481,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"PyramidFlowVAEDecode" : "PyramidFlow VAE Decode",
|
||||
"PyramidFlowTextEncode": "PyramidFlow Text Encode",
|
||||
"PyramidFlowVAEEncode": "PyramidFlow VAE Encode",
|
||||
"PyramidFlowTextEncodeComfy": "PyramidFlow Text Encode Comfy",
|
||||
"PyramidFlowTorchCompileSettings": "PyramidFlow Torch Compile Settings",
|
||||
"PyramidFlowTransformerLoader": "PyramidFlow Model Loader",
|
||||
"PyramidFlowVAELoader": "PyramidFlow VAE Loader"
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user