From 65747a8a7eaebf070ef2d367140c2e96b012e825 Mon Sep 17 00:00:00 2001 From: smthemex <138738845+smthemex@users.noreply.github.com> Date: Wed, 22 Oct 2025 20:34:35 +0800 Subject: [PATCH] Update infer_flashvsr_full.py --- .../examples/WanVSR/infer_flashvsr_full.py | 146 +++++++++++++++--- 1 file changed, 125 insertions(+), 21 deletions(-) diff --git a/FlashVSR/examples/WanVSR/infer_flashvsr_full.py b/FlashVSR/examples/WanVSR/infer_flashvsr_full.py index 7e70dcf..8f49fa8 100644 --- a/FlashVSR/examples/WanVSR/infer_flashvsr_full.py +++ b/FlashVSR/examples/WanVSR/infer_flashvsr_full.py @@ -8,16 +8,34 @@ import imageio from tqdm import tqdm import torch from einops import rearrange - +import folder_paths from ...diffsynth import ModelManager, FlashVSRFullPipeline from .utils.utils import Buffer_LQ4x_Proj from comfy.utils import common_upscale -def tensor2video(frames: torch.Tensor): +def tensor2video(frames): frames = rearrange(frames, "C T H W -> T H W C") - frames = ((frames.float() + 1) * 127.5).clip(0, 255).cpu().numpy().astype(np.uint8) - frames = [Image.fromarray(frame) for frame in frames] - return frames + 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))] @@ -30,6 +48,8 @@ def list_images_natural(folder: str): def largest_8n1_leq(n): # 8n+1 return 0 if n < 1 else ((n - 1)//8)*8 + 1 +def dup_first_frame_1cthw_simple(video_tensor): + return torch.cat([video_tensor[:, :, :1], video_tensor], dim=2) def is_video(path): return os.path.isfile(path) and path.lower().endswith(('.mp4','.mov','.avi','.mkv')) @@ -37,7 +57,7 @@ def is_video(path): 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) + return t.to(dtype=dtype) def save_video(frames, save_path, fps=30, quality=5): os.makedirs(os.path.dirname(save_path), exist_ok=True) @@ -71,19 +91,59 @@ def upscale_then_center_crop(img: Image.Image, scale: int, tW: int, tH: int) -> l = max(0, (sW - tW) // 2); t = max(0, (sH - tH) // 2) return up.crop((l, t, l + tW, t + tH)) -def prepare_input_tensor(path, scale: int = 4,fps=30, dtype=torch.bfloat16, device='cuda'): +def split_into_segments(total_frames, segment_length=81,): + """ + 将总帧数分割为指定长度的段 + + Args: + total_frames: 总帧数 (x) + segment_length: 每段长度 (默认81) + + Returns: + segments: 包含每段起始和结束索引的列表 + """ + segments = [] + start = 0 + while start < total_frames: + end = min(start + segment_length, total_frames) + segments.append((start, end)) + start = end # 无重叠,直接跳到下一组 + + return segments + +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, scale: int = 4,fps=30,dtype=torch.bfloat16, device='cuda'): if isinstance(path,torch.Tensor): total,h0,w0,_ = path.shape sW, sH, tW, tH = compute_scaled_and_target_dims(w0, h0, scale=scale, multiple=128) - path=tensor_upscale(path, tW, tH) + pil_list=tensor2pillist(path) idx = list(range(total)) + [total - 1] * 4 F = largest_8n1_leq(len(idx)) idx = idx[:F] - path=path[idx, :, :, :] - vid=path.permute(3,0,1,2).unsqueeze(0).to(device,dtype=torch.bfloat16) # 1 C F H W - #print(vid.shape) #torch.Size([1, 3, 121, 768, 1280]) + 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 vid, tH, tW, F, fps - + elif os.path.isdir(path): paths0 = list_images_natural(path) if not paths0: @@ -177,7 +237,7 @@ def prepare_input_tensor(path, scale: int = 4,fps=30, dtype=torch.bfloat16, devi else: raise ValueError(f"Unsupported input: {path}") -def init_pipeline(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",device="cuda"): +def init_pipeline(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",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,]) @@ -188,19 +248,23 @@ def init_pipeline(LQ_proj_in_path="./FlashVSR/LQ_proj_in.ckpt",ckpt_path: str = pipe.denoising_model().LQ_proj_in.to(device) pipe.vae.model.encoder = None pipe.vae.model.conv1 = None + 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,prompt_path,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,dtype=torch.bfloat16,device="cuda"): - 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"]) - +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) - video = pipe( + 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), @@ -208,7 +272,47 @@ def run_inference(pipe,prompt_path,input,seed,scale,kv_ratio=3.0,local_range=9,s local_range=local_range, # Recommended: 9 or 11. local_range=9 → sharper details; 11 → more stable results. color_fix = color_fix, ) - #video = tensor2video(video) - #save_video(video, os.path.join(RESULT_ROOT, f"FlashVSR_Full_{name.split('.')[0]}_seed{seed}.mp4"), fps=fps, quality=6) + 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)} + try: + frames = pipe.decode_video(frames, **tiler_kwargs) + except: + pipe.vae.to('cpu') + torch.cuda.empty_cache() + 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_video=dup_first_frame_1cthw_simple(LQ_video) + 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.") - return rearrange(video, "C T H W -> T H W C") + 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