Update infer_flashvsr_full.py

This commit is contained in:
smthemex
2025-10-22 20:34:35 +08:00
committed by GitHub
parent 5c642cf00c
commit 65747a8a7e
+125 -21
View File
@@ -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