Merge pull request #5 from smthemex/Pr1

init
This commit is contained in:
smthemex
2025-10-23 14:47:19 +08:00
committed by GitHub
4 changed files with 77 additions and 74 deletions
+35 -33
View File
@@ -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])
+36 -35
View File
@@ -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
+3 -3
View File
@@ -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
+3 -3
View File
@@ -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"