Files

460 lines
16 KiB
Python
Executable File

import torch
import os
import sys
import gc
script_directory = os.path.dirname(os.path.abspath(__file__))
sys.path.append(script_directory)
from opendit.core.pab_mgr import set_pab_manager
#from opendit.core.parallel_mgr import enable_sequence_parallel, set_parallel_manager
from opendit.models.opensora import RFLOW, OpenSoraVAE_V1_2, STDiT3_XL_2, T5Encoder, text_preprocessing
from opendit.models.opensora.inference_utils import (
append_score_to_prompts,
extract_prompts_loop,
merge_prompt,
prepare_multi_resolution_info,
split_prompt,
apply_mask_strategy
)
import comfy.model_management as mm
from comfy.utils import ProgressBar, load_torch_file
import folder_paths
try:
from flash_attn import flash_attn_varlen_func
FLASH_ATTN_AVAILABLE = True
print("Flash Attention is available")
except:
FLASH_ATTN_AVAILABLE = False
print("WARNING! Flash Attention is not available, using much slower torch SDP attention")
class DownloadAndLoadOpenSoraModel:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": (
[
'hpcai-tech/OpenSora-STDiT-v3'
],
),
"precision": (['fp16','bf16','fp32'],
{
"default": 'bf16'
}),
},
}
RETURN_TYPES = ("OPENDITMODEL",)
RETURN_NAMES = ("opendit_model",)
FUNCTION = "loadmodel"
CATEGORY = "OpenDitWrapper"
def loadmodel(self, model, precision):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_name = model.rsplit('/', 1)[-1]
model_path = os.path.join(folder_paths.models_dir, "opensora", model_name)
if not os.path.exists(model_path):
print(f"Downloading OpenSora model to: {model_path}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id=model,
ignore_patterns=['*ema*'],
local_dir=model_path,
local_dir_use_symlinks=False)
if not hasattr(self, "model"):
print("Loading STDiT...")
self.model = (
STDiT3_XL_2(
from_pretrained=model_path,
qk_norm=True,
enable_flash_attn=FLASH_ATTN_AVAILABLE,
enable_layernorm_kernel=True,
#input_size=latent_size,
in_channels=4,
caption_channels=4096,
model_max_length=300
).to(offload_device, dtype).eval()
)
mm.soft_empty_cache()
opendit_model = {
'model': self.model,
'dtype': dtype
}
return (opendit_model,)
class DownloadAndLoadOpenSoraVAE:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": (
[
'hpcai-tech/OpenSora-VAE-v1.2'
],
),
"precision": (['fp16','bf16','fp32'],
{
"default": 'bf16'
}),
},
}
RETURN_TYPES = ("VAE",)
RETURN_NAMES = ("opendit_vae",)
FUNCTION = "loadmodel"
CATEGORY = "OpenDitWrapper"
def loadmodel(self, model, precision):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_name = model.rsplit('/', 1)[-1]
model_path = os.path.join(folder_paths.models_dir, "opensora", model_name)
if not os.path.exists(model_path):
print(f"Downloading OpenSora model to: {model_path}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id=model,
ignore_patterns=['*ema*'],
local_dir=model_path,
local_dir_use_symlinks=False)
if not hasattr(self, "vae"):
print("Loading VAE...")
self.vae = (
OpenSoraVAE_V1_2(
from_pretrained="hpcai-tech/OpenSora-VAE-v1.2",
micro_frame_size=17,
micro_batch_size=4,
).to(offload_device, dtype).eval()
)
mm.soft_empty_cache()
opendit_model = {
'model': self.vae,
'dtype': dtype
}
return (opendit_model,)
class DownloadAndLoadOpenDiTT5Model:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": (
[
'city96/t5-v1_1-xxl-encoder-bf16'
],
),
"precision": (['fp16','bf16','fp32'],
{
"default": 'bf16'
}),
},
}
RETURN_TYPES = ("OPENDITT5",)
RETURN_NAMES = ("opendit_t5_encoder",)
FUNCTION = "loadmodel"
CATEGORY = "OpenDitWrapper"
def loadmodel(self, model, precision):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_name = model.rsplit('/', 1)[-1]
model_path = os.path.join(folder_paths.models_dir, "t5", model_name)
if not os.path.exists(model_path):
print(f"Downloading OpenSora model to: {model_path}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id=model,
ignore_patterns=['*ema*'],
local_dir=model_path,
local_dir_use_symlinks=False)
if not hasattr(self, "text_encoder"):
print("Loading Text Encoder...")
self.text_encoder = T5Encoder(
from_pretrained=model_path, model_max_length=300, device=device, dtype=dtype, shardformer=False
)
mm.soft_empty_cache()
return (self.text_encoder,)
class OpenDiTConditioning:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"opendit_t5_encoder": ("OPENDITT5",),
"prompt": ("STRING", {"default": "", "multiline": True}),
"camera_prompt": ("STRING", {"default": "", "multiline": True}),
"aesthetic_score": ("FLOAT", {"default": 6.5, "min": 0.0, "max": 100.0, "step": 0.1}),
"flow_score": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100.0, "step": 0.1}),
"keep_model_loaded": ("BOOLEAN", {"default": False}),
},
"optional": {
"opendit_ref": ("OPENDITREF",),
},
}
RETURN_TYPES = ("OPENDITCOND",)
RETURN_NAMES =("opendit_cond",)
FUNCTION = "process"
CATEGORY = "OpenDiTWrapper"
def process(self, opendit_t5_encoder, prompt, camera_prompt, aesthetic_score, flow_score, keep_model_loaded=False, opendit_ref=None):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
self.text_encoder = opendit_t5_encoder
print("process prompt step by step...")
# == process prompt step by step ==
# 0. split prompt
prompt_segment_list, loop_idx_list = split_prompt(prompt)
# 1. append score
prompt_segment_list = append_score_to_prompts(
prompt_segment_list,
aes=aesthetic_score if aesthetic_score > 0 else None,
flow=flow_score if flow_score > 0 else None,
camera_motion=camera_prompt if camera_prompt != "" else None,
)
# 2. clean prompt with T5
prompt_segment_list = [text_preprocessing(prompt) for prompt in prompt_segment_list]
# 3. merge to obtain the final prompt
final_prompt = merge_prompt(prompt_segment_list, loop_idx_list)
final_prompt_loop = extract_prompts_loop([final_prompt], 0)
print("final_prompt_loop: ", final_prompt_loop)
self.text_encoder.t5.model.to(device)
encoded_prompt = self.text_encoder.encode(final_prompt_loop)
if not keep_model_loaded:
self.text_encoder.t5.model.to(offload_device)
mm.soft_empty_cache()
gc.collect()
opendit_cond = {
"encoded_prompt": encoded_prompt,
"refs_x": opendit_ref['refs_x'] if opendit_ref is not None else None,
"mask_strategy": opendit_ref['mask_strategy'] if opendit_ref is not None else None
}
return (opendit_cond,)
class OpenDiTSampler:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"opendit_model": ("OPENDITMODEL",),
"opendit_vae": ("VAE",),
"opendit_cond": ("OPENDITCOND",),
"num_frames": ("INT", {"default": 24, "min": 1, "max": 200, "step": 1}),
"width": ("INT", {"default": 426, "min": 1, "max": 2048, "step": 1}),
"height": ("INT", {"default": 240, "min": 1, "max": 2048, "step": 1}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"steps": ("INT", {"default": 25, "min": 1, "max": 200, "step": 1}),
"cfg": ("FLOAT", {"default": 7.0, "min": 0.0, "max": 20.0, "step": 0.01}),
"fps": ("INT", {"default": 24, "min": 1, "max": 60, "step": 1}),
"keep_model_loaded": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("LATENT", "VAE",)
RETURN_NAMES =("samples", "opendit_vae",)
FUNCTION = "process"
CATEGORY = "OpenDiTWrapper"
def process(self, opendit_model, opendit_vae, opendit_cond, num_frames, width, height, seed, steps, cfg, fps, keep_model_loaded=False):
device = mm.get_torch_device()
dtype = opendit_model['dtype']
offload_device = mm.unet_offload_device()
self.model = opendit_model['model']
set_pab_manager(
steps=steps,
cross_broadcast=True,
cross_threshold=[540, 940],
cross_gap=6,
spatial_broadcast=True,
spatial_threshold=[540, 940],
spatial_gap=2,
temporal_broadcast=True,
temporal_threshold=[540, 940],
temporal_gap=4,
#diffusion_skip=6
#diffusion_skip_timestep= [1,1,1,0,0,0,0,0,0,0]
)
image_size = (height, width)
input_size = (num_frames, *image_size)
latent_size = opendit_vae['model'].get_latent_size(input_size)
scheduler = RFLOW(use_timestep_transform=True, num_sampling_steps=steps, cfg_scale=cfg)
print("Sampling...")
# == sampling ==
torch.manual_seed(seed)
z = torch.randn(1, 4, *latent_size, device=device, dtype=dtype)
mm.soft_empty_cache()
gc.collect()
multi_resolution = "STDiT2"
additional_args = prepare_multi_resolution_info(
multi_resolution, 1, image_size, num_frames, fps, device, dtype
)
print("additional_args: ", additional_args)
final_cond = opendit_cond['encoded_prompt'].copy()
final_cond.update(additional_args)
self.model.to(device)
y_null = self.model.y_embedder.y_embedding[None].repeat(1, 1, 1)[:, None]
final_cond["y"] = torch.cat([final_cond["y"], y_null], 0)
if opendit_cond['refs_x'] is not None:
masks = apply_mask_strategy(z, opendit_cond['refs_x'], opendit_cond['mask_strategy'], 0, align=None)
else:
masks = None
samples = scheduler.sample(
self.model,
final_cond,
z=z,
device=device,
progress=True,
additional_args=additional_args,
mask=masks,
)
if not keep_model_loaded:
self.model.to(offload_device)
mm.soft_empty_cache()
gc.collect()
return (samples, opendit_vae,)
class OpenSoraEncodeReference:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"opendit_vae": ("VAE",),
"ref_image": ("IMAGE", ),
"target_frame_start": (['first','last'],
{
"default": 'first'
}),
"edit_rate": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}),
},
}
RETURN_TYPES = ("OPENDITREF",)
RETURN_NAMES =("opendit_ref",)
FUNCTION = "process"
CATEGORY = "OpenDiTWrapper"
def process(self, opendit_vae, ref_image, target_frame_start, edit_rate):
device = mm.get_torch_device()
dtype = opendit_vae['dtype']
offload_device = mm.unet_offload_device()
self.vae = opendit_vae['model']
# Normalize the tensor
mean = torch.tensor([0.5, 0.5, 0.5]).view(1, 1, 1, -1)
std = torch.tensor([0.5, 0.5, 0.5]).view(1, 1, 1, -1)
normalized_image = (ref_image - mean) / std
normalized_image = normalized_image.permute(3, 0, 1, 2).unsqueeze(0).to(device, dtype)
refs_x = []
ref = []
self.vae.to(device)
r_x = self.vae.encode(normalized_image)
self.vae.to(offload_device)
r_x = r_x.squeeze(0)
ref.append(r_x)
refs_x.append(ref)
frame = 0 if target_frame_start == 'first' else -1
mask_strategy = [f"0,0,0,{frame},{len(ref_image)},{edit_rate}"]
print(mask_strategy)
references = {
"refs_x": refs_x,
"mask_strategy": mask_strategy
}
return (references,)
class OpenSoraDecode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"samples": ("LATENT", ),
"opendit_vae": ("VAE",),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES =("images",)
FUNCTION = "decode"
CATEGORY = "OpenDiTWrapper"
def decode(self, samples, opendit_vae):
device = mm.get_torch_device()
dtype = opendit_vae['dtype']
offload_device = mm.unet_offload_device()
self.vae = opendit_vae['model']
self.vae.to(device)
samples = self.vae.decode(samples.to(dtype),num_frames=len(samples))
self.vae.to(offload_device)
samples = samples.squeeze(0).permute(1, 2, 3, 0).float().cpu()
normalized_tensor = torch.clamp(samples, -1, 1)
tensor_min = normalized_tensor.min()
tensor_max = normalized_tensor.max()
normalized_tensor = (samples - tensor_min) / (tensor_max - tensor_min)
normalized_tensor = torch.clamp(normalized_tensor, 0, 1)
return (normalized_tensor,)
NODE_CLASS_MAPPINGS = {
"OpenDiTSampler": OpenDiTSampler,
"OpenDiTConditioning": OpenDiTConditioning,
"DownloadAndLoadOpenSoraModel": DownloadAndLoadOpenSoraModel,
"DownloadAndLoadOpenSoraVAE": DownloadAndLoadOpenSoraVAE,
"DownloadAndLoadOpenDiTT5Model": DownloadAndLoadOpenDiTT5Model,
"OpenSoraEncodeReference": OpenSoraEncodeReference,
"OpenSoraDecode": OpenSoraDecode
}
NODE_DISPLAY_NAME_MAPPINGS = {
"OpenDiTSampler": "OpenDiT Sampler",
"OpenDiTConditioning": "OpenDiT Conditioning",
"DownloadAndLoadOpenSoraModel": "(Down)Load OpenSora Model",
"DownloadAndLoadOpenSoraVAE": "(Down)Load OpenSora VAE",
"DownloadAndLoadOpenDiTT5Model": "(Down)Load OpenDiT T5 Model",
"OpenSoraEncodeReference": "OpenSora Encode Reference",
"OpenSoraDecode": "OpenSora Decode"
}