diff --git a/FlashVSR/examples/WanVSR/infer_flashvsr_full.py b/FlashVSR/examples/WanVSR/infer_flashvsr_full.py index 96f7226..05e83b4 100644 --- a/FlashVSR/examples/WanVSR/infer_flashvsr_full.py +++ b/FlashVSR/examples/WanVSR/infer_flashvsr_full.py @@ -276,39 +276,41 @@ def run_inference(pipe,input,seed,scale,kv_ratio=3.0,local_range=9,step=1,cfg_sc 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.") - pipe.vae.to('cpu') + with torch.no_grad(): + try: + frames = pipe.decode_video(frames, **tiler_kwargs) + except: + print("vae decode_video OOM.try split latent" ) + 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.") + pipe.vae.to('cpu') del LQ torch.cuda.empty_cache() frames = tensor2video(frames[0]) diff --git a/FlashVSR/examples/WanVSR/infer_flashvsr_tiny.py b/FlashVSR/examples/WanVSR/infer_flashvsr_tiny.py index 9bedd1d..a06829d 100644 --- a/FlashVSR/examples/WanVSR/infer_flashvsr_tiny.py +++ b/FlashVSR/examples/WanVSR/infer_flashvsr_tiny.py @@ -267,46 +267,47 @@ def run_inference_tiny(pipe,input,seed,scale,kv_ratio=3.0,local_range=9,step=1,c pipe.dit.to('cpu') torch.cuda.empty_cache() #print(LQ.shape,frames.shape,LQ_cur_idx)#torch.Size([1, 3, 81, 384, 640]) torch.Size([1, 16, 20, 48, 80]) 77 - try: - frames = pipe.TCDecoder.decode_video(frames.transpose(1, 2),parallel=False, show_progress_bar=False, cond=LQ[:,:,:LQ_cur_idx,:,:]).transpose(1, 2).mul_(2).sub_(1) - except: - print("TCDecoder decode_video OOM.") - pipe.TCDecoder.to('cpu') - torch.cuda.empty_cache() - pipe.TCDecoder.to('cuda') - frames = frames.transpose(1, 2) # 转换为 [B, T, C, H, W] 格式 - cond_frames = LQ[:,:,:LQ_cur_idx,:,:] - segment_size = (split_num-1) * 2 # 160 - latent_segment_size = segment_size // 4 # 40 - decoded_frames_list = [] - total_latent_frames = frames.shape[1] - for start_idx in range(0, total_latent_frames, latent_segment_size): - end_idx = min(start_idx + latent_segment_size, total_latent_frames) + with torch.no_grad(): + try: + frames = pipe.TCDecoder.decode_video(frames.transpose(1, 2),parallel=False, show_progress_bar=False, cond=LQ[:,:,:LQ_cur_idx,:,:]).transpose(1, 2).mul_(2).sub_(1) + except: + print("TCDecoder decode_video OOM.try split latent" ) + pipe.TCDecoder.to('cpu') + torch.cuda.empty_cache() + pipe.TCDecoder.to('cuda') + frames = frames.transpose(1, 2) # 转换为 [B, T, C, H, W] 格式 + cond_frames = LQ[:,:,:LQ_cur_idx,:,:] + segment_size = (split_num-1) * 2 # 160 + latent_segment_size = segment_size // 4 # 40 + decoded_frames_list = [] + total_latent_frames = frames.shape[1] + for start_idx in range(0, total_latent_frames, latent_segment_size): + end_idx = min(start_idx + latent_segment_size, total_latent_frames) + + # 潜在空间的切片 + frames_segment = frames[:, start_idx:end_idx, :, :, :] + + # 对应的条件帧(在原始空间中) + start_cond_idx = start_idx * 4 + end_cond_idx = min(end_idx * 4 , LQ_cur_idx) + cond_segment = cond_frames[:, :, start_cond_idx:end_cond_idx, :, :] + + decoded_segment = pipe.TCDecoder.decode_video( + frames_segment, + parallel=False, + show_progress_bar=False, + cond=cond_segment + ) + decoded_frames_list.append(decoded_segment) - # 潜在空间的切片 - frames_segment = frames[:, start_idx:end_idx, :, :, :] - - # 对应的条件帧(在原始空间中) - start_cond_idx = start_idx * 4 - end_cond_idx = min(end_idx * 4 , LQ_cur_idx) - cond_segment = cond_frames[:, :, start_cond_idx:end_cond_idx, :, :] - - decoded_segment = pipe.TCDecoder.decode_video( - frames_segment, - parallel=False, - show_progress_bar=False, - cond=cond_segment - ) - decoded_frames_list.append(decoded_segment) - - decoded_frames = torch.cat(decoded_frames_list, dim=1) # - frames = decoded_frames.transpose(1, 2).mul_(2).sub_(1) + decoded_frames = torch.cat(decoded_frames_list, dim=1) # + frames = decoded_frames.transpose(1, 2).mul_(2).sub_(1) # 颜色校正(wavelet) # shape: 1,16, 20, 64, 96 try: if color_fix: - if pad_first_frame: + if pad_first_frame: # 加帧 frames = dup_first_frame_1cthw_simple(frames) LQ=dup_first_frame_1cthw_simple(LQ) frames = pipe.ColorCorrector( @@ -316,7 +317,7 @@ def run_inference_tiny(pipe,input,seed,scale,kv_ratio=3.0,local_range=9,step=1,c chunk_size=16, method=fix_method ) - if pad_first_frame: + if pad_first_frame: #减帧 frames = frames[:, :, 1:, :, :] # remove first frame except: pass diff --git a/FlashVSR/requirements.txt b/FlashVSR/requirements.txt index 99c216f..8f22762 100644 --- a/FlashVSR/requirements.txt +++ b/FlashVSR/requirements.txt @@ -8,7 +8,7 @@ einops==0.8.1 huggingface-hub==0.34.4 matplotlib==3.10.3 numpy==1.26.4 -opencv-python==4.11.0.86 +opencv-python==4.11.0.86 # opencv-python-headless==4.11.0.86 peft==0.16.0 pillow==11.0.0 @@ -16,10 +16,10 @@ safetensors==0.5.3 sentencepiece==0.2.0 transformers==4.46.2 pytorch-lightning==2.5.2 -imageio==2.37.0 +imageio==2.37.0 # imageio-ffmpeg==0.6.0 protobuf==3.20.3 -ftfy==6.3.1 +ftfy==6.3.1 # pandas==2.3.0 tqdm datasets \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index 1a79bb2..7acc8b6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,9 +1,9 @@ [project] name = "flashvsr" -description = "" -version = "1.0.0" +description = "FlashVSR:Towards Real-Time Diffusion-Based Streaming Video Super-Resolution,you can use it in comfyUI" +version = "1.0.1" license = {file = "LICENSE"} -dependencies = ["torch", "torchaudio", "torchmetrics", "torchsde", "torchvision", "accelerate", "einops", "huggingface-hub", "matplotlib", "numpy", "opencv-python", "opencv-python-headless", "peft", "pillow0", "safetensors", "sentencepiece", "transformers", "pytorch-lightning", "imageio", "imageio-ffmpeg", "protobuf", "ftfy", "pandas", "tqdm", "datasets"] +dependencies = ["torch", "torchaudio", "torchmetrics", "torchsde", "torchvision", "accelerate", "einops", "huggingface-hub", "matplotlib", "numpy", "opencv-python", "opencv-python-headless", "peft", "pillow", "safetensors", "sentencepiece", "transformers", "pytorch-lightning", "imageio", "imageio-ffmpeg", "protobuf", "ftfy", "pandas", "tqdm", "datasets"] [project.urls] Repository = "https://github.com/smthemex/ComfyUI_FlashVSR"