This commit is contained in:
smthemex
2025-11-05 22:03:19 +08:00
parent 93676f0728
commit 375d3245e0
94 changed files with 1097 additions and 5 deletions
Binary file not shown.
Binary file not shown.
@@ -165,6 +165,7 @@ class FlashVSRFullPipeline(BasePipeline):
self.ColorCorrector = TorchColorCorrectorWavelet(levels=5)
self.new_decoder=False
self.VAE=None
self.version="1.0"
print(r"""
@@ -164,7 +164,7 @@ class FlashVSRTinyPipeline(BasePipeline):
self.prompt_emb_posi = None
self.ColorCorrector = TorchColorCorrectorWavelet(levels=5)
self.long_mode=False
self.version="1.0"
print(r"""
███████╗██╗ █████╗ ███████╗██╗ ██╗██╗ ██╗███████╗█████╗
@@ -164,6 +164,7 @@ class FlashVSRTinyLongPipeline(BasePipeline):
self.prompt_emb_posi = None
self.ColorCorrector = TorchColorCorrectorWavelet(levels=5)
self.long_mode=True
self.version="1.0"
print(r"""
███████╗██╗ █████╗ ███████╗██╗ ██╗██╗ ██╗███████╗█████╗
██╔════╝██║ ██╔══██╗██╔════╝██║ ██║██║ ██║██╔════╝██╔══██╗
@@ -0,0 +1,401 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import os, re, time
import numpy as np
from PIL import Image
import imageio
from tqdm import tqdm
import torch
from einops import rearrange
from safetensors.torch import load_file
from ...diffsynth import ModelManager, FlashVSRFullPipeline
from .utils.utils import Causal_LQ4x_Proj
import folder_paths
import torch.nn.functional as F
def tensor2video(frames: torch.Tensor):
frames = rearrange(frames, "C T H W -> T H W C")
try:
frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8)
frames = [Image.fromarray(frame) for frame in frames]
return frames
except:
try:
frames=frames.cpu()
frames = ((frames.float() + 1) * 127.5).clip(0, 255).numpy().astype(np.uint8)
frames = [Image.fromarray(frame) for frame in frames]
return frames
except:
batch_size = min(32, frames.shape[0])
total_frames = frames.shape[0]
frame_list = []
for i in range(0, total_frames, batch_size):
batch_frames = frames[i:min(i + batch_size, total_frames)]
batch_frames = ((batch_frames.float() + 1) * 127.5).clip(0, 255)
batch_frames_np = batch_frames.cpu().numpy().astype(np.uint8)
for frame in batch_frames_np:
frame_list.append(Image.fromarray(frame))
return frame_list
def natural_key(name: str):
return [int(t) if t.isdigit() else t.lower() for t in re.split(r'([0-9]+)', os.path.basename(name))]
def list_images_natural(folder: str):
exts = ('.png', '.jpg', '.jpeg', '.PNG', '.JPG', '.JPEG')
fs = [os.path.join(folder, f) for f in os.listdir(folder) if f.endswith(exts)]
fs.sort(key=natural_key)
return fs
def largest_8n1_leq(n): # 8n+1
return 0 if n < 1 else ((n - 1)//8)*8 + 1
def is_video(path):
return os.path.isfile(path) and path.lower().endswith(('.mp4','.mov','.avi','.mkv'))
def pil_to_tensor_neg1_1(img: Image.Image, dtype=torch.bfloat16, device='cuda'):
t = torch.from_numpy(np.asarray(img, np.uint8)).to(device=device, dtype=torch.float32) # HWC
t = t.permute(2,0,1) / 255.0 * 2.0 - 1.0 # CHW in [-1,1]
return t.to(dtype)
def save_video(frames, save_path, fps=30, quality=5):
os.makedirs(os.path.dirname(save_path), exist_ok=True)
w = imageio.get_writer(save_path, fps=fps, quality=quality)
for f in tqdm(frames, desc=f"Saving {os.path.basename(save_path)}"):
w.append_data(np.array(f))
w.close()
def compute_scaled_and_target_dims(w0: int, h0: int, scale: int = 4, multiple: int = 128):
if w0 <= 0 or h0 <= 0:
raise ValueError("invalid original size")
sW, sH = w0 * scale, h0 * scale
tW = max(multiple, (sW // multiple) * multiple)
tH = max(multiple, (sH // multiple) * multiple)
return sW, sH, tW, tH
def upscale_then_center_crop(img: Image.Image, scale: int, tW: int, tH: int) -> Image.Image:
w0, h0 = img.size
sW, sH = w0 * scale, h0 * scale
# 先放大
up = img.resize((sW, sH), Image.BICUBIC)
# 中心裁剪
l = max(0, (sW - tW) // 2); t = max(0, (sH - tH) // 2)
return up.crop((l, t, l + tW, t + tH))
def tensor2image(tensor):
tensor = tensor.cpu()
image_np = tensor.squeeze().mul(255).clamp(0, 255).byte().numpy()
image = Image.fromarray(image_np, mode='RGB')
return image
def dup_first_frame_1cthw_simple(video_tensor):
return torch.cat([video_tensor[:, :, :1], video_tensor], dim=2)
def tensor2pillist(tensor_in):
d1, _, _, _ = tensor_in.size()
if d1 == 1:
img_list = [tensor2image(tensor_in)]
else:
tensor_list = torch.chunk(tensor_in, chunks=d1)
img_list=[tensor2image(i) for i in tensor_list]
return img_list
def prepare_input_tensor(path: str, scale: int = 4,fps=30, dtype=torch.bfloat16, device='cuda'):
if isinstance(path,torch.Tensor):
total,h0,w0,_ = path.shape
if total == 1:
print("got image,repeating to 25 frames")
path = path.repeat(25, 1, 1, 1)
total=25
sW, sH, tW, tH = compute_scaled_and_target_dims(w0, h0, scale=scale, multiple=128)
pil_list=tensor2pillist(path)
idx = list(range(total)) + [total - 1] * 4
F = largest_8n1_leq(len(idx))
idx = idx[:F]
frames = []
pil_list = [pil_list[i] for i in idx]
for i in idx:
img = pil_list[i].convert('RGB')
img_out = upscale_then_center_crop(img, scale=scale, tW=tW, tH=tH)
frames.append(pil_to_tensor_neg1_1(img_out, dtype, device))
frames = torch.stack(frames, 0).permute(1,0,2,3).unsqueeze(0) # 1 C F H W
torch.cuda.empty_cache()
return frames, tH, tW, F, fps
elif os.path.isdir(path):
paths0 = list_images_natural(path)
if not paths0:
raise FileNotFoundError(f"No images in {path}")
with Image.open(paths0[0]) as _img0:
w0, h0 = _img0.size
N0 = len(paths0)
print(f"[{os.path.basename(path)}] Original Resolution: {w0}x{h0} | Original Frames: {N0}")
sW, sH, tW, tH = compute_scaled_and_target_dims(w0, h0, scale=scale, multiple=128)
print(f"[{os.path.basename(path)}] Scaled Resolution (x{scale}): {sW}x{sH} -> Target (128-multiple): {tW}x{tH}")
paths = paths0 + [paths0[-1]] * 4
F = largest_8n1_leq(len(paths))
if F == 0:
raise RuntimeError(f"Not enough frames after padding in {path}. Got {len(paths)}.")
paths = paths[:F]
print(f"[{os.path.basename(path)}] Target Frames (8n-3): {F-4}")
frames = []
for p in paths:
with Image.open(p).convert('RGB') as img:
img_out = upscale_then_center_crop(img, scale=scale, tW=tW, tH=tH)
frames.append(pil_to_tensor_neg1_1(img_out, dtype, device))
vid = torch.stack(frames, 0).permute(1,0,2,3).unsqueeze(0)
fps = 30
return vid, tH, tW, F, fps
elif is_video(path):
rdr = imageio.get_reader(path)
first = Image.fromarray(rdr.get_data(0)).convert('RGB')
w0, h0 = first.size
meta = {}
try:
meta = rdr.get_meta_data()
except Exception:
pass
fps_val = meta.get('fps', 30)
fps = int(round(fps_val)) if isinstance(fps_val, (int, float)) else 30
def count_frames(r):
try:
nf = meta.get('nframes', None)
if isinstance(nf, int) and nf > 0:
return nf
except Exception:
pass
try:
return r.count_frames()
except Exception:
n = 0
try:
while True:
r.get_data(n); n += 1
except Exception:
return n
total = count_frames(rdr)
if total <= 0:
rdr.close()
raise RuntimeError(f"Cannot read frames from {path}")
print(f"[{os.path.basename(path)}] Original Resolution: {w0}x{h0} | Original Frames: {total} | FPS: {fps}")
sW, sH, tW, tH = compute_scaled_and_target_dims(w0, h0, scale=scale, multiple=128)
print(f"[{os.path.basename(path)}] Scaled Resolution (x{scale}): {sW}x{sH} -> Target (128-multiple): {tW}x{tH}")
idx = list(range(total)) + [total - 1] * 4
F = largest_8n1_leq(len(idx))
if F == 0:
rdr.close()
raise RuntimeError(f"Not enough frames after padding in {path}. Got {len(idx)}.")
idx = idx[:F]
print(f"[{os.path.basename(path)}] Target Frames (8n-3): {F-4}")
frames = []
try:
for i in idx:
img = Image.fromarray(rdr.get_data(i)).convert('RGB')
img_out = upscale_then_center_crop(img, scale=scale, tW=tW, tH=tH)
frames.append(pil_to_tensor_neg1_1(img_out, dtype, device))
finally:
try:
rdr.close()
except Exception:
pass
vid = torch.stack(frames, 0).permute(1,0,2,3).unsqueeze(0) # 1 C F H W
return vid, tH, tW, F, fps
else:
raise ValueError(f"Unsupported input: {path}")
def init_pipeline_v11(prompt_path,LQ_proj_in_path="./FlashVSR/LQ_proj_in.ckpt",ckpt_path: str = "./FlashVSR/diffusion_pytorch_model_streaming_dmd.safetensors", vae_path: str = "./FlashVSR/Wan2.1_VAE.pth",decode_vae="none",cur_dir="",device="cuda"):
#print(torch.cuda.current_device(), torch.cuda.get_device_name(torch.cuda.current_device()))
mm = ModelManager(torch_dtype=torch.bfloat16, device="cpu")
mm.load_models([ckpt_path,vae_path,])
new_decoder=True if decode_vae!="none" else False
if new_decoder:
pipe.new_decoder = True
if "light" in decode_vae.lower() or "tae" in decode_vae.lower():
if os.path.basename(decode_vae).split(".")[0]=="lightvaew2_1":
from ...vae import WanVAE
print("use lightvae decoder")
VAE = WanVAE(vae_path=decode_vae,dtype=torch.bfloat16,device=device,use_lightvae=True)
elif os.path.basename(decode_vae).split(".")[0]=="taew2_1":
from ...vae_tiny import WanVAE_tiny
print("use vae_tiny decoder")
VAE = WanVAE_tiny(vae_path=vae_path,dtype=torch.bfloat16,device=device,need_scaled=False)
elif os.path.basename(decode_vae).split(".")[0]=="lighttaew2_1":
from ...vae_tiny import WanVAE_tiny
print("use vae_tiny light decoder")
VAE = WanVAE_tiny(vae_path=decode_vae,dtype=torch.bfloat16,device=device,need_scaled=True)
else:
raise ValueError(f"Unknown vae_name: {decode_vae},only support lightvae,tae,tae_tiny,lighttae_tiny")
pipe.VAE=VAE
else:
print("use upscale2x decoder")
from diffusers import AutoencoderKLWan
config=AutoencoderKLWan.load_config(os.path.join(cur_dir,"FlashVSR/examples/config.json"))
VAE=AutoencoderKLWan.from_config(config).to(device,dtype=torch.bfloat16)
vae_dict=load_file(decode_vae,device="cpu")
VAE.load_state_dict(vae_dict,strict=False)
pipe.VAE=VAE
del vae_dict
pipe = FlashVSRFullPipeline.from_model_manager(mm, device="cuda")
pipe.denoising_model().LQ_proj_in = Causal_LQ4x_Proj(in_dim=3, out_dim=1536, layer_num=1).to("cuda", dtype=torch.bfloat16)
#LQ_proj_in_path = "./FlashVSR-v1.1/LQ_proj_in.ckpt"
if os.path.exists(LQ_proj_in_path):
pipe.denoising_model().LQ_proj_in.load_state_dict(torch.load(LQ_proj_in_path, map_location="cpu",weights_only=False), strict=True)
pipe.denoising_model().LQ_proj_in.to('cuda')
pipe.vae.model.encoder = None
pipe.vae.model.conv1 = None
#pipe.to('cuda');
pipe.enable_vram_management(num_persistent_param_in_dit=None)
pipe.init_cross_kv(prompt_path); pipe.load_models_to_device(["dit","vae"])
return pipe
def run_inference(pipe,input,seed,scale,kv_ratio=3.0,local_range=9,step=1,cfg_scale=1.0,sparse_ratio=2.0,tiled=True,color_fix=True,fix_method="wavelet",split_num=81,dtype=torch.bfloat16,device="cuda",save_vodeo_=False,):
pipe.to('cuda') #pipe.enable_vram_management(num_persistent_param_in_dit=None)
#pipe.init_cross_kv(prompt_path); pipe.load_models_to_device(["dit","vae"])
pad_first_frame = True if "wavelet"== fix_method and color_fix else False
#total,h0,w0,_ = input.shape
torch.cuda.empty_cache(); torch.cuda.ipc_collect()
LQ, th, tw, F, fps = prepare_input_tensor(input, scale=scale, dtype=dtype, device=device)
frames = pipe(
prompt="", negative_prompt="", cfg_scale=cfg_scale, num_inference_steps=step, seed=seed, tiled=tiled,
LQ_video=LQ, num_frames=F, height=th, width=tw, is_full_block=False, if_buffer=True,
topk_ratio=sparse_ratio*768*1280/(th*tw),
kv_ratio=kv_ratio,
local_range=local_range, # Recommended: 9 or 11. local_range=9 → sharper details; 11 → more stable results.
color_fix = color_fix,
)
pipe.dit.to('cpu')
torch.cuda.empty_cache()
#torch.Size([1, 16, 20, 48, 80])
tiler_kwargs = {"tiled": tiled, "tile_size": (60, 104), "tile_stride": (30, 52)}
with torch.no_grad():
try:
frames = pipe.decode_video(frames, **tiler_kwargs)
except:
print("vae decode_video OOM.try split latent" )
if pipe.new_decoder:
if pipe.VAE.__class__.__name__ == "AutoencoderKLWan":
pipe.VAE.to('cpu')
else:
if pipe.VAE.__class__.__name__ == "WanVAE":
pipe.VAE.to_cpu()
else: pass
else:
pipe.vae.to('cpu')
torch.cuda.empty_cache()
if pipe.new_decoder:
if pipe.VAE.__class__.__name__ == "AutoencoderKLWan":
pipe.VAE.to('cuda')
else:
if pipe.VAE.__class__.__name__ == "WanVAE":
pipe.VAE.to_cuda()
else: pass
else:
pipe.vae.to('cuda')
total_frames = frames.shape[2]
segment_size = (split_num-1) * 2 // 4 # 40
decoded_frames_list = []
for start_idx in range(0, total_frames, segment_size):
end_idx = min(start_idx + segment_size, total_frames)
frames_segment = frames[:, :, start_idx:end_idx, :, :]
decoded_segment = pipe.decode_video(frames_segment, **tiler_kwargs)
decoded_frames_list.append(decoded_segment)
frames = torch.cat(decoded_frames_list, dim=2)
try:
if color_fix:
if pad_first_frame:
frames = dup_first_frame_1cthw_simple(frames)
LQ=dup_first_frame_1cthw_simple(LQ)
if pipe.new_decoder and LQ.shape[-1]!=frames.shape[-1]:
scale_=int(frames.shape[-1]/LQ.shape[-1])
LQ=upscale_lq_video_bilinear(LQ,scale_)
frames = pipe.ColorCorrector(
frames.to(device=device),
LQ[:, :, :frames.shape[2], :, :],
clip_range=(-1, 1),
chunk_size=16,
method=fix_method
)
if pad_first_frame:
frames = frames[:, :, 1:, :, :] # remove first frame
except:
pass
print("Done.")
pipe.vae.to('cpu')
del LQ
torch.cuda.empty_cache()
frames = tensor2video(frames[0])
if save_vodeo_:
save_video(frames, os.path.join(folder_paths.get_output_directory(),f"FlashVSR_Full_seed{seed}.mp4"), fps=fps, quality=6)
return frames
def upscale_lq_video_bilinear(LQ_video,scale_):
B, C, T, H, W = LQ_video.shape
LQ_reshaped = LQ_video.view(B*T, C, H, W)
HQ_reshaped = F.interpolate(
LQ_reshaped,
size=(H*scale_, W*scale_),
mode='bilinear',
align_corners=False
)
HQ_video = HQ_reshaped.view(B, C, T, H*scale_, W*scale_)
return HQ_video
# def main():
# RESULT_ROOT = "./results"
# os.makedirs(RESULT_ROOT, exist_ok=True)
# inputs = [
# "./inputs/example0.mp4",
# "./inputs/example1.mp4",
# "./inputs/example2.mp4",
# "./inputs/example3.mp4",
# ]
# seed, scale, dtype, device = 0, 4, torch.bfloat16, 'cuda'
# sparse_ratio = 2.0 # Recommended: 1.5 or 2.0. 1.5 → faster; 2.0 → more stable.
# pipe = init_pipeline()
# for p in inputs:
# torch.cuda.empty_cache(); torch.cuda.ipc_collect()
# name = os.path.basename(p.rstrip('/'))
# if name.startswith('.'):
# continue
# try:
# LQ, th, tw, F, fps = prepare_input_tensor(p, scale=scale, dtype=dtype, device=device)
# except Exception as e:
# print(f"[Error] {name}: {e}")
# continue
# video = pipe(
# prompt="", negative_prompt="", cfg_scale=1.0, num_inference_steps=1, seed=seed,
# tiled=False,# Disable tiling: faster inference but higher VRAM usage.
# # Set to True for lower memory consumption at the cost of speed.
# LQ_video=LQ, num_frames=F, height=th, width=tw, is_full_block=False, if_buffer=True,
# topk_ratio=sparse_ratio*768*1280/(th*tw),
# kv_ratio=3.0,
# local_range=11, # Recommended: 9 or 11. local_range=9 → sharper details; 11 → more stable results.
# color_fix = True,
# )
# video = tensor2video(video)
# save_video(video, os.path.join(RESULT_ROOT, f"FlashVSR_v1.1_Full_{name.split('.')[0]}_seed{seed}.mp4"), fps=fps, quality=6)
# print("Done.")
# if __name__ == "__main__":
# main()
@@ -0,0 +1,287 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import os, re, time
import numpy as np
from PIL import Image
import imageio
from tqdm import tqdm
import torch
from einops import rearrange
import folder_paths
from ...diffsynth import ModelManager, FlashVSRTinyPipeline
from .utils.utils import Causal_LQ4x_Proj
from .utils.TCDecoder import build_tcdecoder
def tensor2video(frames):
frames = rearrange(frames, "C T H W -> T H W C")
try:
frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8)
frames = [Image.fromarray(frame) for frame in frames]
return frames
except:
try:
frames=frames.cpu()
frames = ((frames.float() + 1) * 127.5).clip(0, 255).numpy().astype(np.uint8)
frames = [Image.fromarray(frame) for frame in frames]
return frames
except:
batch_size = min(32, frames.shape[0])
total_frames = frames.shape[0]
frame_list = []
for i in range(0, total_frames, batch_size):
batch_frames = frames[i:min(i + batch_size, total_frames)]
batch_frames = ((batch_frames.float() + 1) * 127.5).clip(0, 255)
batch_frames_np = batch_frames.cpu().numpy().astype(np.uint8)
for frame in batch_frames_np:
frame_list.append(Image.fromarray(frame))
return frame_list
def natural_key(name: str):
return [int(t) if t.isdigit() else t.lower() for t in re.split(r'([0-9]+)', os.path.basename(name))]
def list_images_natural(folder: str):
exts = ('.png', '.jpg', '.jpeg', '.PNG', '.JPG', '.JPEG')
fs = [os.path.join(folder, f) for f in os.listdir(folder) if f.endswith(exts)]
fs.sort(key=natural_key)
return fs
def largest_8n1_leq(n): # 8n+1
return 0 if n < 1 else ((n - 1)//8)*8 + 1
def is_video(path):
return os.path.isfile(path) and path.lower().endswith(('.mp4','.mov','.avi','.mkv'))
def pil_to_tensor_neg1_1(img: Image.Image, dtype=torch.bfloat16, device='cuda'):
t = torch.from_numpy(np.asarray(img, np.uint8)).to(device=device, dtype=torch.float32) # HWC
t = t.permute(2,0,1) / 255.0 * 2.0 - 1.0 # CHW in [-1,1]
return t.to(dtype)
def save_video(frames, save_path, fps=30, quality=5):
os.makedirs(os.path.dirname(save_path), exist_ok=True)
w = imageio.get_writer(save_path, fps=fps, quality=quality)
for f in tqdm(frames, desc=f"Saving {os.path.basename(save_path)}"):
w.append_data(np.array(f))
w.close()
def compute_scaled_and_target_dims(w0: int, h0: int, scale: float = 4.0, multiple: int = 128):
if w0 <= 0 or h0 <= 0:
raise ValueError("Invalid original size")
if scale <= 0:
raise ValueError("scale must be > 0")
sW = int(round(w0 * scale))
sH = int(round(h0 * scale))
tW = (sW // multiple) * multiple
tH = (sH // multiple) * multiple
if tW == 0 or tH == 0:
raise ValueError(
f"Scaled size too small ({sW}x{sH}) for multiple={multiple}. "
f"Increase scale (got {scale})."
)
return sW, sH, tW, tH
def upscale_then_center_crop(img: Image.Image, scale: float, tW: int, tH: int) -> Image.Image:
w0, h0 = img.size
sW = int(round(w0 * scale))
sH = int(round(h0 * scale))
if tW > sW or tH > sH:
raise ValueError(
f"Target crop ({tW}x{tH}) exceeds scaled size ({sW}x{sH}). "
f"Increase scale."
)
up = img.resize((sW, sH), Image.BICUBIC)
l = (sW - tW) // 2
t = (sH - tH) // 2
return up.crop((l, t, l + tW, t + tH))
def tensor2image(tensor):
tensor = tensor.cpu()
image_np = tensor.squeeze().mul(255).clamp(0, 255).byte().numpy()
image = Image.fromarray(image_np, mode='RGB')
return image
def tensor2pillist(tensor_in):
d1, _, _, _ = tensor_in.size()
if d1 == 1:
img_list = [tensor2image(tensor_in)]
else:
tensor_list = torch.chunk(tensor_in, chunks=d1)
img_list=[tensor2image(i) for i in tensor_list]
return img_list
def prepare_input_tensor(path: str, scale: float = 4, fps=30,dtype=torch.bfloat16, device='cuda'):
if isinstance(path,torch.Tensor):
total,h0,w0,_ = path.shape
if total == 1:
print("got image,repeating to 25 frames")
path = path.repeat(25, 1, 1, 1)
total=25
sW, sH, tW, tH = compute_scaled_and_target_dims(w0, h0, scale=scale, multiple=128)
pil_list=tensor2pillist(path)
idx = list(range(total)) + [total - 1] * 4
F = largest_8n1_leq(len(idx))
idx = idx[:F]
frames = []
pil_list = [pil_list[i] for i in idx]
for i in idx:
img = pil_list[i].convert('RGB')
img_out = upscale_then_center_crop(img, scale=scale, tW=tW, tH=tH)
frames.append(pil_to_tensor_neg1_1(img_out, dtype, device))
frames = torch.stack(frames, 0).permute(1,0,2,3).unsqueeze(0) # 1 C F H W
torch.cuda.empty_cache()
return frames, tH, tW, F, fps
if os.path.isdir(path):
paths0 = list_images_natural(path)
if not paths0:
raise FileNotFoundError(f"No images in {path}")
with Image.open(paths0[0]) as _img0:
w0, h0 = _img0.size
N0 = len(paths0)
print(f"[{os.path.basename(path)}] Original Resolution: {w0}x{h0} | Original Frames: {N0}")
sW, sH, tW, tH = compute_scaled_and_target_dims(w0, h0, scale=scale, multiple=128)
print(f"[{os.path.basename(path)}] Scaled (x{scale:.2f}): {sW}x{sH} -> Target (128-multiple): {tW}x{tH}")
paths = paths0 + [paths0[-1]] * 4
F = largest_8n1_leq(len(paths))
if F == 0:
raise RuntimeError(f"Not enough frames after padding in {path}. Got {len(paths)}.")
paths = paths[:F]
print(f"[{os.path.basename(path)}] Target Frames (8n-3): {F-4}")
frames = []
for p in paths:
with Image.open(p).convert('RGB') as img:
img_out = upscale_then_center_crop(img, scale=scale, tW=tW, tH=tH)
frames.append(pil_to_tensor_neg1_1(img_out, dtype, device))
vid = torch.stack(frames, 0).permute(1,0,2,3).unsqueeze(0) # 1 C F H W
fps = 30
return vid, tH, tW, F, fps
if is_video(path):
rdr = imageio.get_reader(path)
first = Image.fromarray(rdr.get_data(0)).convert('RGB')
w0, h0 = first.size
meta = {}
try: meta = rdr.get_meta_data()
except Exception: pass
fps_val = meta.get('fps', 30)
fps = int(round(fps_val)) if isinstance(fps_val, (int, float)) else 30
def count_frames(r):
try:
nf = meta.get('nframes', None)
if isinstance(nf,int) and nf>0: return nf
except Exception: pass
try: return r.count_frames()
except Exception:
n=0
try:
while True: r.get_data(n); n+=1
except Exception:
return n
total = count_frames(rdr)
if total <= 0:
rdr.close()
raise RuntimeError(f"Cannot read frames from {path}")
print(f"[{os.path.basename(path)}] Original Resolution: {w0}x{h0} | Original Frames: {total} | FPS: {fps}")
sW, sH, tW, tH = compute_scaled_and_target_dims(w0, h0, scale=scale, multiple=128)
print(f"[{os.path.basename(path)}] Scaled (x{scale:.2f}): {sW}x{sH} -> Target (128-multiple): {tW}x{tH}")
idx = list(range(total)) + [total-1]*4
F = largest_8n1_leq(len(idx))
if F == 0:
rdr.close()
raise RuntimeError(f"Not enough frames after padding in {path}. Got {len(idx)}.")
idx = idx[:F]
print(f"[{os.path.basename(path)}] Target Frames (8n-3): {F-4}")
frames = []
try:
for i in idx:
img = Image.fromarray(rdr.get_data(i)).convert('RGB')
img_out = upscale_then_center_crop(img, scale=scale, tW=tW, tH=tH)
frames.append(pil_to_tensor_neg1_1(img_out, dtype, device))
finally:
try: rdr.close()
except Exception: pass
vid = torch.stack(frames, 0).permute(1,0,2,3).unsqueeze(0) # 1 C F H W
return vid, tH, tW, F, fps
raise ValueError(f"Unsupported input: {path}")
def init_pipeline_v11_tiny(prompt_path,LQ_proj_in_path = "./FlashVSR/LQ_proj_in.ckpt",ckpt_path="./FlashVSR/diffusion_pytorch_model_streaming_dmd.safetensors",TCDecoder_path="./FlashVSR/TCDecoder.ckpt",device="cuda"):
#print(torch.cuda.current_device(), torch.cuda.get_device_name(torch.cuda.current_device()))
mm = ModelManager(torch_dtype=torch.bfloat16, device="cpu")
mm.load_models([ckpt_path,])
pipe = FlashVSRTinyPipeline.from_model_manager(mm, device="cuda")
pipe.denoising_model().LQ_proj_in = Causal_LQ4x_Proj(in_dim=3, out_dim=1536, layer_num=1).to("cuda", dtype=torch.bfloat16)
#LQ_proj_in_path = "./FlashVSR-v1.1/LQ_proj_in.ckpt"
if os.path.exists(LQ_proj_in_path):
pipe.denoising_model().LQ_proj_in.load_state_dict(torch.load(LQ_proj_in_path, map_location="cpu"), strict=True)
pipe.denoising_model().LQ_proj_in.to('cuda')
multi_scale_channels = [512, 256, 128, 128]
pipe.TCDecoder = build_tcdecoder(new_channels=multi_scale_channels, new_latent_channels=16+768)
mis = pipe.TCDecoder.load_state_dict(torch.load(TCDecoder_path,weights_only=False,), strict=False)
print(mis)
#pipe.to('cuda');
pipe.enable_vram_management(num_persistent_param_in_dit=None)
pipe.init_cross_kv(prompt_path); pipe.load_models_to_device(["dit","vae"])
return pipe
# def main():
# RESULT_ROOT = "./results"
# os.makedirs(RESULT_ROOT, exist_ok=True)
# inputs = [
# "./inputs/example0.mp4",
# "./inputs/example1.mp4",
# "./inputs/example2.mp4",
# "./inputs/example3.mp4",
# ]
# seed, scale, dtype, device = 0, 4.0, torch.bfloat16, 'cuda'
# sparse_ratio = 2.0 # Recommended: 1.5 or 2.0. 1.5 → faster; 2.0 → more stable.
# pipe = init_pipeline()
# for p in inputs:
# torch.cuda.empty_cache(); torch.cuda.ipc_collect()
# name = os.path.basename(p.rstrip('/'))
# if name.startswith('.'):
# continue
# try:
# LQ, th, tw, F, fps = prepare_input_tensor(p, scale=scale, dtype=dtype, device=device)
# except Exception as e:
# print(f"[Error] {name}: {e}"); continue
# video = pipe(
# prompt="", negative_prompt="", cfg_scale=1.0, num_inference_steps=1, seed=seed,
# LQ_video=LQ, num_frames=F, height=th, width=tw, is_full_block=False, if_buffer=True,
# topk_ratio=sparse_ratio*768*1280/(th*tw),
# kv_ratio=3.0,
# local_range=11, # Recommended: 9 or 11. local_range=9 → sharper details; 11 → more stable results.
# color_fix = True,
# )
# video = tensor2video(video)
# save_video(video, os.path.join(RESULT_ROOT, f"FlashVSR_v1.1_Tiny_{name.split('.')[0]}_seed{seed}.mp4"), fps=fps, quality=6)
# print("Done.")
# if __name__ == "__main__":
# main()
@@ -0,0 +1,288 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import os, re, time
import numpy as np
from PIL import Image
import imageio
from tqdm import tqdm
import torch
from einops import rearrange
import folder_paths
from ...diffsynth import ModelManager, FlashVSRTinyLongPipeline
from .utils.utils import Causal_LQ4x_Proj
from .utils.TCDecoder import build_tcdecoder
def tensor2video(frames):
frames = rearrange(frames, "C T H W -> T H W C")
try:
frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8)
frames = [Image.fromarray(frame) for frame in frames]
return frames
except:
try:
frames=frames.cpu()
frames = ((frames.float() + 1) * 127.5).clip(0, 255).numpy().astype(np.uint8)
frames = [Image.fromarray(frame) for frame in frames]
return frames
except:
batch_size = min(32, frames.shape[0])
total_frames = frames.shape[0]
frame_list = []
for i in range(0, total_frames, batch_size):
batch_frames = frames[i:min(i + batch_size, total_frames)]
batch_frames = ((batch_frames.float() + 1) * 127.5).clip(0, 255)
batch_frames_np = batch_frames.cpu().numpy().astype(np.uint8)
for frame in batch_frames_np:
frame_list.append(Image.fromarray(frame))
return frame_list
def natural_key(name: str):
return [int(t) if t.isdigit() else t.lower() for t in re.split(r'([0-9]+)', os.path.basename(name))]
def list_images_natural(folder: str):
exts = ('.png', '.jpg', '.jpeg', '.PNG', '.JPG', '.JPEG')
fs = [os.path.join(folder, f) for f in os.listdir(folder) if f.endswith(exts)]
fs.sort(key=natural_key)
return fs
def largest_8n1_leq(n): # 8n+1
return 0 if n < 1 else ((n - 1)//8)*8 + 1
def is_video(path):
return os.path.isfile(path) and path.lower().endswith(('.mp4','.mov','.avi','.mkv'))
def pil_to_tensor_neg1_1(img: Image.Image, dtype=torch.bfloat16, device='cuda'):
t = torch.from_numpy(np.asarray(img, np.uint8)).to(device=device, dtype=torch.float32) # HWC
t = t.permute(2,0,1) / 255.0 * 2.0 - 1.0 # CHW in [-1,1]
return t.to(dtype)
def save_video(frames, save_path, fps=30, quality=5):
os.makedirs(os.path.dirname(save_path), exist_ok=True)
w = imageio.get_writer(save_path, fps=fps, quality=quality)
for f in tqdm(frames, desc=f"Saving {os.path.basename(save_path)}"):
w.append_data(np.array(f))
w.close()
def compute_scaled_and_target_dims(w0: int, h0: int, scale: float = 4.0, multiple: int = 128):
if w0 <= 0 or h0 <= 0:
raise ValueError("Invalid original size")
if scale <= 0:
raise ValueError("scale must be > 0")
sW = int(round(w0 * scale))
sH = int(round(h0 * scale))
tW = (sW // multiple) * multiple
tH = (sH // multiple) * multiple
if tW == 0 or tH == 0:
raise ValueError(
f"Scaled size too small ({sW}x{sH}) for multiple={multiple}. "
f"Increase scale (got {scale})."
)
return sW, sH, tW, tH
def upscale_then_center_crop(img: Image.Image, scale: float, tW: int, tH: int) -> Image.Image:
w0, h0 = img.size
sW = int(round(w0 * scale))
sH = int(round(h0 * scale))
if tW > sW or tH > sH:
raise ValueError(
f"Target crop ({tW}x{tH}) exceeds scaled size ({sW}x{sH}). "
f"Increase scale."
)
up = img.resize((sW, sH), Image.BICUBIC)
l = (sW - tW) // 2
t = (sH - tH) // 2
return up.crop((l, t, l + tW, t + tH))
def tensor2image(tensor):
tensor = tensor.cpu()
image_np = tensor.squeeze().mul(255).clamp(0, 255).byte().numpy()
image = Image.fromarray(image_np, mode='RGB')
return image
def tensor2pillist(tensor_in):
d1, _, _, _ = tensor_in.size()
if d1 == 1:
img_list = [tensor2image(tensor_in)]
else:
tensor_list = torch.chunk(tensor_in, chunks=d1)
img_list=[tensor2image(i) for i in tensor_list]
return img_list
def prepare_input_tensor(path: str, scale: float = 4,fps=30, dtype=torch.bfloat16, device='cuda'):
if isinstance(path,torch.Tensor):
total,h0,w0,_ = path.shape
if total == 1:
print("got image,repeating to 25 frames")
path = path.repeat(25, 1, 1, 1)
total=25
sW, sH, tW, tH = compute_scaled_and_target_dims(w0, h0, scale=scale, multiple=128)
pil_list=tensor2pillist(path)
idx = list(range(total)) + [total - 1] * 4
F = largest_8n1_leq(len(idx))
idx = idx[:F]
frames = []
pil_list = [pil_list[i] for i in idx]
for i in idx:
img = pil_list[i].convert('RGB')
img_out = upscale_then_center_crop(img, scale=scale, tW=tW, tH=tH)
frames.append(pil_to_tensor_neg1_1(img_out, dtype, device))
frames = torch.stack(frames, 0).permute(1,0,2,3).unsqueeze(0) # 1 C F H W
torch.cuda.empty_cache()
return frames, tH, tW, F, fps
if os.path.isdir(path):
paths0 = list_images_natural(path)
if not paths0:
raise FileNotFoundError(f"No images in {path}")
with Image.open(paths0[0]) as _img0:
w0, h0 = _img0.size
N0 = len(paths0)
print(f"[{os.path.basename(path)}] Original Resolution: {w0}x{h0} | Original Frames: {N0}")
sW, sH, tW, tH = compute_scaled_and_target_dims(w0, h0, scale=scale, multiple=128)
print(f"[{os.path.basename(path)}] Scaled (x{scale:.2f}): {sW}x{sH} -> Target (128-multiple): {tW}x{tH}")
paths = paths0 + [paths0[-1]] * 4
F = largest_8n1_leq(len(paths))
if F == 0:
raise RuntimeError(f"Not enough frames after padding in {path}. Got {len(paths)}.")
paths = paths[:F]
print(f"[{os.path.basename(path)}] Target Frames (8n-3): {F-4}")
frames = []
count_img = 0
for p in paths:
with Image.open(p).convert('RGB') as img:
img_out = upscale_then_center_crop(img, scale=scale, tW=tW, tH=tH)
frames.append(pil_to_tensor_neg1_1(img_out, dtype, 'cpu'))
print(count_img, len(paths), end = '\r')
count_img+=1
vid = torch.stack(frames, 0).permute(1,0,2,3).unsqueeze(0) # 1 C F H W
fps = 30
return vid, tH, tW, F, fps
if is_video(path):
rdr = imageio.get_reader(path)
first = Image.fromarray(rdr.get_data(0)).convert('RGB')
w0, h0 = first.size
meta = {}
try: meta = rdr.get_meta_data()
except Exception: pass
fps_val = meta.get('fps', 30)
fps = int(round(fps_val)) if isinstance(fps_val, (int, float)) else 30
def count_frames(r):
try:
nf = meta.get('nframes', None)
if isinstance(nf,int) and nf>0: return nf
except Exception: pass
try: return r.count_frames()
except Exception:
n=0
try:
while True: r.get_data(n); n+=1
except Exception:
return n
total = count_frames(rdr)
if total <= 0:
rdr.close()
raise RuntimeError(f"Cannot read frames from {path}")
print(f"[{os.path.basename(path)}] Original Resolution: {w0}x{h0} | Original Frames: {total} | FPS: {fps}")
sW, sH, tW, tH = compute_scaled_and_target_dims(w0, h0, scale=scale, multiple=128)
print(f"[{os.path.basename(path)}] Scaled (x{scale:.2f}): {sW}x{sH} -> Target (128-multiple): {tW}x{tH}")
idx = list(range(total)) + [total-1]*4
F = largest_8n1_leq(len(idx))
if F == 0:
rdr.close()
raise RuntimeError(f"Not enough frames after padding in {path}. Got {len(idx)}.")
idx = idx[:F]
print(f"[{os.path.basename(path)}] Target Frames (8n-3): {F-4}")
frames = []
try:
for i in idx:
img = Image.fromarray(rdr.get_data(i)).convert('RGB')
img_out = upscale_then_center_crop(img, scale=scale, tW=tW, tH=tH)
frames.append(pil_to_tensor_neg1_1(img_out, dtype, 'cpu'))
print(i, len(idx), end = '\r')
finally:
try: rdr.close()
except Exception: pass
vid = torch.stack(frames, 0).permute(1,0,2,3).unsqueeze(0) # 1 C F H W
return vid, tH, tW, F, fps
raise ValueError(f"Unsupported input: {path}")
def init_pipeline_long_v11(prompt_path,LQ_proj_in_path = "./FlashVSR/LQ_proj_in.ckpt",ckpt_path="./FlashVSR/diffusion_pytorch_model_streaming_dmd.safetensors",TCDecoder_path="./FlashVSR/TCDecoder.ckpt",device="cuda"):
#print(torch.cuda.current_device(), torch.cuda.get_device_name(torch.cuda.current_device()))
mm = ModelManager(torch_dtype=torch.bfloat16, device="cpu")
mm.load_models([ckpt_path,])
pipe = FlashVSRTinyLongPipeline.from_model_manager(mm, device="cuda")
pipe.denoising_model().LQ_proj_in = Causal_LQ4x_Proj(in_dim=3, out_dim=1536, layer_num=1).to("cuda", dtype=torch.bfloat16)
#LQ_proj_in_path = "./FlashVSR-v1.1/LQ_proj_in.ckpt"
if os.path.exists(LQ_proj_in_path):
pipe.denoising_model().LQ_proj_in.load_state_dict(torch.load(LQ_proj_in_path, map_location="cpu",weights_only=False,), strict=True)
pipe.denoising_model().LQ_proj_in.to('cuda')
multi_scale_channels = [512, 256, 128, 128]
pipe.TCDecoder = build_tcdecoder(new_channels=multi_scale_channels, new_latent_channels=16+768)
mis = pipe.TCDecoder.load_state_dict(torch.load(TCDecoder_path,weights_only=False,), strict=False)
print(mis)
#pipe.to('cuda');
pipe.enable_vram_management(num_persistent_param_in_dit=None)
pipe.init_cross_kv(prompt_path); pipe.load_models_to_device(["dit","vae"])
return pipe
# def main():
# RESULT_ROOT = "./results"
# os.makedirs(RESULT_ROOT, exist_ok=True)
# inputs = [
# "./inputs/example4.mp4",
# ]
# seed, scale, dtype, device = 0, 4.0, torch.bfloat16, 'cuda'
# sparse_ratio = 2.0 # Recommended: 1.5 or 2.0. 1.5 → faster; 2.0 → more stable.
# pipe = init_pipeline()
# for p in inputs:
# torch.cuda.empty_cache(); torch.cuda.ipc_collect()
# name = os.path.basename(p.rstrip('/'))
# if name.startswith('.'):
# continue
# try:
# LQ, th, tw, F, fps = prepare_input_tensor(p, scale=scale, dtype=dtype, device=device)
# except Exception as e:
# print(f"[Error] {name}: {e}"); continue
# video = pipe(
# prompt="", negative_prompt="", cfg_scale=1.0, num_inference_steps=1, seed=seed,
# LQ_video=LQ, num_frames=F, height=th, width=tw, is_full_block=False, if_buffer=True,
# topk_ratio=sparse_ratio*768*1280/(th*tw),
# kv_ratio=3.0,
# local_range=11, # Recommended: 9 or 11. local_range=9 → sharper details; 11 → more stable results.
# color_fix = True,
# )
# video = tensor2video(video)
# save_video(video, os.path.join(RESULT_ROOT, f"FlashVSR_v1.1_Tiny_Long_{name.split('.')[0]}_seed{seed}.mp4"), fps=fps, quality=5)
# print("Done.")
# if __name__ == "__main__":
# main()
+100
View File
@@ -272,3 +272,103 @@ class Buffer_LQ4x_Proj(nn.Module):
outputs.append(self.linear_layers[i](out_x))
self.clip_idx += 1
return outputs
class Causal_LQ4x_Proj(nn.Module):
def __init__(self, in_dim, out_dim, layer_num=30):
super().__init__()
self.ff = 1
self.hh = 16
self.ww = 16
self.hidden_dim1 = 2048
self.hidden_dim2 = 3072
self.layer_num = layer_num
self.pixel_shuffle = PixelShuffle3d(self.ff, self.hh, self.ww)
self.conv1 = CausalConv3d(in_dim*self.ff*self.hh*self.ww, self.hidden_dim1, (4, 3, 3), stride=(2, 1, 1), padding=(1, 1, 1)) # f -> f/2 h -> h w -> w
self.norm1 = RMS_norm(self.hidden_dim1, images=False)
self.act1 = nn.SiLU()
self.conv2 = CausalConv3d(self.hidden_dim1, self.hidden_dim2, (4, 3, 3), stride=(2, 1, 1), padding=(1, 1, 1)) # f -> f/2 h -> h w -> w
self.norm2 = RMS_norm(self.hidden_dim2, images=False)
self.act2 = nn.SiLU()
self.linear_layers = nn.ModuleList([nn.Linear(self.hidden_dim2, out_dim) for _ in range(layer_num)])
self.clip_idx = 0
def forward(self, video):
self.clear_cache()
# x: (B, C, F, H, W)
t = video.shape[2]
iter_ = 1 + (t - 1) // 4
first_frame = video[:, :, :1, :, :].repeat(1, 1, 3, 1, 1)
video = torch.cat([first_frame, video], dim=2)
# print(video.shape)
out_x = []
for i in range(iter_):
x = self.pixel_shuffle(video[:,:,i*4:(i+1)*4,:,:])
cache1_x = x[:, :, -CACHE_T:, :, :].clone()
x = self.conv1(x, self.cache['conv1'])
self.cache['conv1'] = cache1_x
x = self.norm1(x)
x = self.act1(x)
cache2_x = x[:, :, -CACHE_T:, :, :].clone()
if i == 0:
self.cache['conv2'] = cache2_x
continue
x = self.conv2(x, self.cache['conv2'])
self.cache['conv2'] = cache2_x
x = self.norm2(x)
x = self.act2(x)
out_x.append(x)
out_x = torch.cat(out_x, dim = 2)
out_x = rearrange(out_x, 'b c f h w -> b (f h w) c')
outputs = []
for i in range(self.layer_num):
outputs.append(self.linear_layers[i](out_x))
return outputs
def clear_cache(self):
self.cache = {}
self.cache['conv1'] = None
self.cache['conv2'] = None
self.clip_idx = 0
def stream_forward(self, video_clip):
if self.clip_idx == 0:
# self.clear_cache()
first_frame = video_clip[:, :, :1, :, :].repeat(1, 1, 3, 1, 1)
video_clip = torch.cat([first_frame, video_clip], dim=2)
x = self.pixel_shuffle(video_clip)
cache1_x = x[:, :, -CACHE_T:, :, :].clone()
x = self.conv1(x, self.cache['conv1'])
self.cache['conv1'] = cache1_x
x = self.norm1(x)
x = self.act1(x)
cache2_x = x[:, :, -CACHE_T:, :, :].clone()
self.cache['conv2'] = cache2_x
self.clip_idx += 1
return None
else:
x = self.pixel_shuffle(video_clip)
cache1_x = x[:, :, -CACHE_T:, :, :].clone()
x = self.conv1(x, self.cache['conv1'])
self.cache['conv1'] = cache1_x
x = self.norm1(x)
x = self.act1(x)
cache2_x = x[:, :, -CACHE_T:, :, :].clone()
x = self.conv2(x, self.cache['conv2'])
self.cache['conv2'] = cache2_x
x = self.norm2(x)
x = self.act2(x)
out_x = rearrange(x, 'b c f h w -> b (f h w) c')
outputs = []
for i in range(self.layer_num):
outputs.append(self.linear_layers[i](out_x))
self.clip_idx += 1
return outputs
+18 -4
View File
@@ -8,6 +8,9 @@ from .model_loader_utils import tensor_upscale,load_images_list,get_video_files
from .FlashVSR.examples.WanVSR.infer_flashvsr_full import init_pipeline,run_inference
from .FlashVSR.examples.WanVSR.infer_flashvsr_tiny import init_pipeline_tiny,run_inference_tiny
from .FlashVSR.examples.WanVSR.infer_flashvsr_tiny_long_video import init_pipeline_long,run_inference_tiny_long
from .FlashVSR.examples.WanVSR.infer_flashvsr_v11_full import init_pipeline_v11
from .FlashVSR.examples.WanVSR.infer_flashvsr_v11_tiny import init_pipeline_v11_tiny
from .FlashVSR.examples.WanVSR.infer_flashvsr_v11_tiny_long_video import init_pipeline_long_v11
import folder_paths
from typing_extensions import override
from comfy_api.latest import ComfyExtension, io
@@ -46,13 +49,14 @@ class FlashVSR_SM_Model(io.ComfyNode):
io.Combo.Input("tcd_encoder",options= ["none"] + [i for i in folder_paths.get_filename_list("FlashVSR") if "tcd" in i.lower()] ),
io.Boolean.Input("tiny_long", default=False),
io.Combo.Input("decode_vae",options= ["none"] + folder_paths.get_filename_list("vae") ),
io.Combo.Input("version",options= ["1.1","1.0"] ),
],
outputs=[
io.Custom("FlashVSR_SM_Model").Output(),
],
)
@classmethod
def execute(cls, dit,proj_pt,emb_pt,vae,tcd_encoder,tiny_long,decode_vae) -> io.NodeOutput:
def execute(cls, dit,proj_pt,emb_pt,vae,tcd_encoder,tiny_long,decode_vae,version) -> io.NodeOutput:
dit_path=folder_paths.get_full_path("FlashVSR", dit) if dit != "none" else None
proj_pt_path=folder_paths.get_full_path("FlashVSR", proj_pt) if proj_pt != "none" else None
vae_path=folder_paths.get_full_path("vae", vae) if vae != "none" else None
@@ -63,14 +67,24 @@ class FlashVSR_SM_Model(io.ComfyNode):
assert vae_path is not None or tcd_encoder_path is not None , "Please select the Sdit,proj_pt,checkpoint file"
if tcd_encoder_path is not None:
if tiny_long:
model=init_pipeline_long(prompt_path,proj_pt_path,dit_path, tcd_encoder_path, device="cuda")
if "1.0"==version:
model=init_pipeline_long(prompt_path,proj_pt_path,dit_path, tcd_encoder_path, device="cuda")
else:
model=init_pipeline_long_v11(prompt_path,proj_pt_path,dit_path, tcd_encoder_path, device="cuda")
else:
model=init_pipeline_tiny(prompt_path,proj_pt_path,dit_path, tcd_encoder_path, device="cuda")
if "1.0"==version:
model=init_pipeline_tiny(prompt_path,proj_pt_path,dit_path, tcd_encoder_path, device="cuda")
else:
model=init_pipeline_v11_tiny(prompt_path,proj_pt_path,dit_path, tcd_encoder_path, device="cuda")
elif vae_path is not None :
decode_vae=folder_paths.get_full_path("vae", decode_vae) if decode_vae != "none" else "none"
model=init_pipeline(prompt_path,proj_pt_path,dit_path, vae_path,decode_vae,node_cr_path ,device="cuda")
if "1.0"==version:
model=init_pipeline(prompt_path,proj_pt_path,dit_path, vae_path,decode_vae,node_cr_path ,device="cuda")
else:
model=init_pipeline_v11(prompt_path,proj_pt_path,dit_path, vae_path,decode_vae,node_cr_path ,device="cuda")
else:
raise Exception("Please select the vae or tcd_encoder")
model.version = version
return io.NodeOutput(model)
Binary file not shown.
Binary file not shown.
Binary file not shown.