init
This commit is contained in:
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
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"""
|
||||
███████╗██╗ █████╗ ███████╗██╗ ██╗██╗ ██╗███████╗█████╗
|
||||
██╔════╝██║ ██╔══██╗██╔════╝██║ ██║██║ ██║██╔════╝██╔══██╗
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -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()
|
||||
Binary file not shown.
Binary file not shown.
@@ -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
@@ -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.
Reference in New Issue
Block a user