Add files via upload
This commit is contained in:
@@ -18,6 +18,14 @@ from ..models.wan_video_vae import WanVideoVAE, RMS_norm, CausalConv3d, Upsample
|
||||
from ..schedulers.flow_match import FlowMatchScheduler
|
||||
from .base import BasePipeline
|
||||
|
||||
# --- Safetensors support ---
|
||||
st_load_file = None # Define the variable in global scope first
|
||||
try:
|
||||
from safetensors.torch import load_file as st_load_file
|
||||
except ImportError:
|
||||
# st_load_file remains None if import fails
|
||||
print("Warning: 'safetensors' not installed. Safetensors (.safetensors) files cannot be loaded.")
|
||||
# ---------------------------
|
||||
|
||||
# -----------------------------
|
||||
# 基础工具:ADAIN 所需的统计量(保留以备需要;管线默认用 wavelet)
|
||||
@@ -80,7 +88,7 @@ def _wavelet_reconstruct(content: torch.Tensor, style: torch.Tensor, levels: int
|
||||
_, s_low = _wavelet_decompose(style, levels=levels)
|
||||
return c_high + s_low
|
||||
|
||||
|
||||
from ..models.utils import clean_vram
|
||||
# -----------------------------
|
||||
# 无状态颜色矫正模块(视频友好,默认 wavelet)
|
||||
# -----------------------------
|
||||
@@ -165,13 +173,12 @@ class FlashVSRFullPipeline(BasePipeline):
|
||||
self.ColorCorrector = TorchColorCorrectorWavelet(levels=5)
|
||||
|
||||
print(r"""
|
||||
███████╗██╗ █████╗ ███████╗██╗ ██╗██╗ ██╗███████╗█████╗
|
||||
██╔════╝██║ ██╔══██╗██╔════╝██║ ██║██║ ██║██╔════╝██╔══██╗
|
||||
█████╗ ██║ ███████║███████╗███████║╚██╗ ██╔╝███████╗███████║
|
||||
██╔══╝ ██║ ██╔══██║╚════██║██╔══██║ ╚████╔╝ ╚════██║██╔═██║
|
||||
██║ ███████╗██║ ██║███████║██║ ██║ ╚██╔╝ ███████║██║ ██║
|
||||
╚═╝ ╚══════╝╚═╝ ╚═╝╚══════╝╚═╝ ╚═╝ ╚═╝ ╚══════╝╚═╝ ╚═╝
|
||||
⚡FlashVSR
|
||||
███████╗██╗ █████╗ ███████╗██╗ ██╗██╗ ██╗███████╗█████╗
|
||||
██╔════╝██║ ██╔══██╗██╔════╝██║ ██║██║ ██║██╔════╝██╔══██╗
|
||||
█████╗ ██║ ███████║███████╗███████║╚██╗ ██╔╝███████╗███████║
|
||||
██╔══╝ ██║ ██╔══██║╚════██║██╔══██║ ╚████╔╝ ╚════██║██╔═██║
|
||||
██║ ███████╗██║ ██║███████║██║ ██║ ╚██╔╝ ███████║██║ ██║
|
||||
╚═╝ ╚══════╝╚═╝ ╚═╝╚══════╝╚═╝ ╚═╝ ╚═╝ ╚══════╝╚═╝ ╚═╝
|
||||
""")
|
||||
|
||||
def enable_vram_management(self, num_persistent_param_in_dit=None):
|
||||
@@ -250,20 +257,44 @@ class FlashVSRFullPipeline(BasePipeline):
|
||||
def init_cross_kv(
|
||||
self,
|
||||
context_tensor: Optional[torch.Tensor] = None,
|
||||
prompt_path = None
|
||||
):
|
||||
self.load_models_to_device(["dit"])
|
||||
"""
|
||||
使用固定 prompt 生成文本 context,并在 WanModel 中初始化所有 CrossAttention 的 KV 缓存。
|
||||
必须在 __call__ 前显式调用一次。
|
||||
"""
|
||||
prompt_path = "../../examples/WanVSR/prompt_tensor/posi_prompt.pth"
|
||||
#prompt_path = "../../examples/WanVSR/prompt_tensor/posi_prompt.pth"
|
||||
if self.dit is None:
|
||||
raise RuntimeError("请先通过 fetch_models / from_model_manager 初始化 self.dit")
|
||||
|
||||
if context_tensor is None:
|
||||
if prompt_path is None:
|
||||
raise ValueError("init_cross_kv: 需要提供 prompt_path 或 context_tensor 其一")
|
||||
ctx = torch.load(prompt_path, map_location=self.device)
|
||||
|
||||
# --- Safetensors loading logic added here ---
|
||||
prompt_path_lower = prompt_path.lower()
|
||||
if prompt_path_lower.endswith(".safetensors"):
|
||||
if st_load_file is None:
|
||||
raise ImportError("The 'safetensors' library must be installed to load .safetensors files.")
|
||||
|
||||
# Load the tensor from safetensors
|
||||
loaded_dict = st_load_file(prompt_path, device=self.device)
|
||||
|
||||
# Safetensors loads a dict. Assuming the context tensor is the only or primary key.
|
||||
if len(loaded_dict) == 1:
|
||||
ctx = list(loaded_dict.values())[0]
|
||||
elif 'context' in loaded_dict: # Common key for text context
|
||||
ctx = loaded_dict['context']
|
||||
else:
|
||||
raise ValueError(f"Safetensors file {prompt_path} does not contain an obvious single tensor ('context' key not found and multiple keys exist).")
|
||||
|
||||
else:
|
||||
# Default behavior for .pth, .pt, etc.
|
||||
ctx = torch.load(prompt_path, map_location=self.device)
|
||||
|
||||
# --------------------------------------------
|
||||
# ctx = torch.load(prompt_path, map_location=self.device)
|
||||
else:
|
||||
ctx = context_tensor
|
||||
|
||||
@@ -326,6 +357,8 @@ class FlashVSRFullPipeline(BasePipeline):
|
||||
kv_ratio=3.0,
|
||||
local_range = 9,
|
||||
color_fix = True,
|
||||
unload_dit = False,
|
||||
skip_vae = False,
|
||||
):
|
||||
# 只接受 cfg=1.0(与原代码一致)
|
||||
assert cfg_scale == 1.0, "cfg_scale must be 1.0"
|
||||
@@ -361,15 +394,17 @@ class FlashVSRFullPipeline(BasePipeline):
|
||||
if hasattr(self.dit, "LQ_proj_in"):
|
||||
self.dit.LQ_proj_in.clear_cache()
|
||||
|
||||
|
||||
latents_total = []
|
||||
self.vae.clear_cache()
|
||||
|
||||
if unload_dit and hasattr(self, 'dit') and self.dit is not None:
|
||||
current_dit_device = next(iter(self.dit.parameters())).device
|
||||
if str(current_dit_device) != str(self.device):
|
||||
print(f"[FlashVSR] DiT is on {current_dit_device}, moving it to target device {self.device}...")
|
||||
self.dit.to(self.device)
|
||||
|
||||
with torch.no_grad():
|
||||
for cur_process_idx in tqdm(range(process_total_num)):
|
||||
torch.cuda.synchronize()
|
||||
dit_time_start = time.time()
|
||||
|
||||
for cur_process_idx in progress_bar_cmd(range(process_total_num)):
|
||||
if cur_process_idx == 0:
|
||||
pre_cache_k = [None] * len(self.dit.blocks)
|
||||
pre_cache_v = [None] * len(self.dit.blocks)
|
||||
@@ -402,7 +437,7 @@ class FlashVSRFullPipeline(BasePipeline):
|
||||
for layer_idx in range(len(LQ_latents)):
|
||||
LQ_latents[layer_idx] = torch.cat([LQ_latents[layer_idx], cur[layer_idx]], dim=1)
|
||||
cur_latents = latents[:, :, 4+cur_process_idx*2:6+cur_process_idx*2, :, :]
|
||||
|
||||
|
||||
# 推理(无 motion_controller / vace)
|
||||
noise_pred_posi, pre_cache_k, pre_cache_v = model_fn_wan_video(
|
||||
self.dit,
|
||||
@@ -427,10 +462,26 @@ class FlashVSRFullPipeline(BasePipeline):
|
||||
# 更新 latent
|
||||
cur_latents = cur_latents - noise_pred_posi
|
||||
latents_total.append(cur_latents)
|
||||
|
||||
if unload_dit and hasattr(self, 'dit') and not next(self.dit.parameters()).is_cpu:
|
||||
try:
|
||||
del pre_cache_k, pre_cache_v
|
||||
except NameError:
|
||||
pass
|
||||
print("[FlashVSR] Offloading DiT to the CPU to free up VRAM...")
|
||||
self.dit.to('cpu')
|
||||
clean_vram()
|
||||
|
||||
latents = torch.cat(latents_total, dim=2)
|
||||
|
||||
|
||||
del latents_total
|
||||
clean_vram()
|
||||
|
||||
if skip_vae:
|
||||
return latents
|
||||
|
||||
# Decode
|
||||
print("[FlashVSR] Starting VAE decoding...")
|
||||
frames = self.decode_video(latents, **tiler_kwargs)
|
||||
|
||||
# 颜色校正(wavelet)
|
||||
|
||||
@@ -18,6 +18,14 @@ from ..models.wan_video_vae import WanVideoVAE, RMS_norm, CausalConv3d, Upsample
|
||||
from ..schedulers.flow_match import FlowMatchScheduler
|
||||
from .base import BasePipeline
|
||||
|
||||
# --- Safetensors support ---
|
||||
st_load_file = None # Define the variable in global scope first
|
||||
try:
|
||||
from safetensors.torch import load_file as st_load_file
|
||||
except ImportError:
|
||||
# st_load_file remains None if import fails
|
||||
print("Warning: 'safetensors' not installed. Safetensors (.safetensors) files cannot be loaded.")
|
||||
# ---------------------------
|
||||
|
||||
# -----------------------------
|
||||
# 基础工具:ADAIN 所需的统计量(保留以备需要;管线默认用 wavelet)
|
||||
@@ -80,7 +88,7 @@ def _wavelet_reconstruct(content: torch.Tensor, style: torch.Tensor, levels: int
|
||||
_, s_low = _wavelet_decompose(style, levels=levels)
|
||||
return c_high + s_low
|
||||
|
||||
|
||||
from ..models.utils import clean_vram
|
||||
# -----------------------------
|
||||
# 无状态颜色矫正模块(视频友好,默认 wavelet)
|
||||
# -----------------------------
|
||||
@@ -165,13 +173,12 @@ class FlashVSRTinyPipeline(BasePipeline):
|
||||
self.ColorCorrector = TorchColorCorrectorWavelet(levels=5)
|
||||
|
||||
print(r"""
|
||||
███████╗██╗ █████╗ ███████╗██╗ ██╗██╗ ██╗███████╗█████╗
|
||||
██╔════╝██║ ██╔══██╗██╔════╝██║ ██║██║ ██║██╔════╝██╔══██╗
|
||||
█████╗ ██║ ███████║███████╗███████║╚██╗ ██╔╝███████╗███████║
|
||||
██╔══╝ ██║ ██╔══██║╚════██║██╔══██║ ╚████╔╝ ╚════██║██╔═██║
|
||||
██║ ███████╗██║ ██║███████║██║ ██║ ╚██╔╝ ███████║██║ ██║
|
||||
╚═╝ ╚══════╝╚═╝ ╚═╝╚══════╝╚═╝ ╚═╝ ╚═╝ ╚══════╝╚═╝ ╚═╝
|
||||
⚡FlashVSR
|
||||
███████╗██╗ █████╗ ███████╗██╗ ██╗██╗ ██╗███████╗█████╗
|
||||
██╔════╝██║ ██╔══██╗██╔════╝██║ ██║██║ ██║██╔════╝██╔══██╗
|
||||
█████╗ ██║ ███████║███████╗███████║╚██╗ ██╔╝███████╗███████║
|
||||
██╔══╝ ██║ ██╔══██║╚════██║██╔══██║ ╚████╔╝ ╚════██║██╔═██║
|
||||
██║ ███████╗██║ ██║███████║██║ ██║ ╚██╔╝ ███████║██║ ██║
|
||||
╚═╝ ╚══════╝╚═╝ ╚═╝╚══════╝╚═╝ ╚═╝ ╚═╝ ╚══════╝╚═╝ ╚═╝
|
||||
""")
|
||||
|
||||
def enable_vram_management(self, num_persistent_param_in_dit=None):
|
||||
@@ -229,13 +236,14 @@ class FlashVSRTinyPipeline(BasePipeline):
|
||||
def init_cross_kv(
|
||||
self,
|
||||
context_tensor: Optional[torch.Tensor] = None,
|
||||
prompt_path = None,
|
||||
):
|
||||
self.load_models_to_device(["dit"])
|
||||
"""
|
||||
使用固定 prompt 生成文本 context,并在 WanModel 中初始化所有 CrossAttention 的 KV 缓存。
|
||||
必须在 __call__ 前显式调用一次。
|
||||
"""
|
||||
prompt_path = "../../examples/WanVSR/prompt_tensor/posi_prompt.pth"
|
||||
#prompt_path = "../../examples/WanVSR/prompt_tensor/posi_prompt.pth"
|
||||
|
||||
if self.dit is None:
|
||||
raise RuntimeError("请先通过 fetch_models / from_model_manager 初始化 self.dit")
|
||||
@@ -243,7 +251,30 @@ class FlashVSRTinyPipeline(BasePipeline):
|
||||
if context_tensor is None:
|
||||
if prompt_path is None:
|
||||
raise ValueError("init_cross_kv: 需要提供 prompt_path 或 context_tensor 其一")
|
||||
ctx = torch.load(prompt_path, map_location=self.device)
|
||||
|
||||
# --- Safetensors loading logic added here ---
|
||||
prompt_path_lower = prompt_path.lower()
|
||||
if prompt_path_lower.endswith(".safetensors"):
|
||||
if st_load_file is None:
|
||||
raise ImportError("The 'safetensors' library must be installed to load .safetensors files.")
|
||||
|
||||
# Load the tensor from safetensors
|
||||
loaded_dict = st_load_file(prompt_path, device=self.device)
|
||||
|
||||
# Safetensors loads a dict. Assuming the context tensor is the only or primary key.
|
||||
if len(loaded_dict) == 1:
|
||||
ctx = list(loaded_dict.values())[0]
|
||||
elif 'context' in loaded_dict: # Common key for text context
|
||||
ctx = loaded_dict['context']
|
||||
else:
|
||||
raise ValueError(f"Safetensors file {prompt_path} does not contain an obvious single tensor ('context' key not found and multiple keys exist).")
|
||||
|
||||
else:
|
||||
# Default behavior for .pth, .pt, etc.
|
||||
ctx = torch.load(prompt_path, map_location=self.device)
|
||||
|
||||
# --------------------------------------------
|
||||
# ctx = torch.load(prompt_path, map_location=self.device)
|
||||
else:
|
||||
ctx = context_tensor
|
||||
|
||||
@@ -274,9 +305,19 @@ class FlashVSRTinyPipeline(BasePipeline):
|
||||
latents = self.vae.encode(input_video, device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
|
||||
return latents
|
||||
|
||||
def decode_video(self, latents, tiled=True, tile_size=(34, 34), tile_stride=(18, 16)):
|
||||
def _decode_video(self, latents, tiled=True, tile_size=(34, 34), tile_stride=(18, 16)):
|
||||
frames = self.vae.decode(latents, device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
|
||||
return frames
|
||||
|
||||
def decode_video(self, latents, cond=None, **kwargs):
|
||||
frames = self.TCDecoder.decode_video(
|
||||
latents.transpose(1, 2), # TCDecoder 需要 (B, F, C, H, W)
|
||||
parallel=False,
|
||||
show_progress_bar=False,
|
||||
cond=cond
|
||||
).transpose(1, 2).mul_(2).sub_(1) # 转回 (B, C, F, H, W) 格式,范围 -1 to 1
|
||||
|
||||
return frames
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
@@ -306,10 +347,12 @@ class FlashVSRTinyPipeline(BasePipeline):
|
||||
kv_ratio=3.0,
|
||||
local_range = 9,
|
||||
color_fix = True,
|
||||
unload_dit = False,
|
||||
skip_vae = False,
|
||||
):
|
||||
# 只接受 cfg=1.0(与原代码一致)
|
||||
assert cfg_scale == 1.0, "cfg_scale must be 1.0"
|
||||
|
||||
|
||||
# 要求:必须先 init_cross_kv()
|
||||
if self.prompt_emb_posi is None or 'context' not in self.prompt_emb_posi:
|
||||
raise RuntimeError(
|
||||
@@ -347,9 +390,15 @@ class FlashVSRTinyPipeline(BasePipeline):
|
||||
self.TCDecoder.clean_mem()
|
||||
LQ_pre_idx = 0
|
||||
LQ_cur_idx = 0
|
||||
|
||||
if unload_dit and hasattr(self, 'dit') and self.dit is not None:
|
||||
current_dit_device = next(iter(self.dit.parameters())).device
|
||||
if str(current_dit_device) != str(self.device):
|
||||
print(f"[FlashVSR] DiT is on {current_dit_device}, moving it to target device {self.device}...")
|
||||
self.dit.to(self.device)
|
||||
|
||||
with torch.no_grad():
|
||||
for cur_process_idx in tqdm(range(process_total_num)):
|
||||
for cur_process_idx in progress_bar_cmd(range(process_total_num)):
|
||||
if cur_process_idx == 0:
|
||||
pre_cache_k = [None] * len(self.dit.blocks)
|
||||
pre_cache_v = [None] * len(self.dit.blocks)
|
||||
@@ -384,7 +433,7 @@ class FlashVSRTinyPipeline(BasePipeline):
|
||||
LQ_latents[layer_idx] = torch.cat([LQ_latents[layer_idx], cur[layer_idx]], dim=1)
|
||||
LQ_cur_idx = cur_process_idx*8+21+(inner_loop_num-2)*4
|
||||
cur_latents = latents[:, :, 4+cur_process_idx*2:6+cur_process_idx*2, :, :]
|
||||
|
||||
|
||||
# 推理(无 motion_controller / vace)
|
||||
noise_pred_posi, pre_cache_k, pre_cache_v = model_fn_wan_video(
|
||||
self.dit,
|
||||
@@ -410,10 +459,26 @@ class FlashVSRTinyPipeline(BasePipeline):
|
||||
cur_latents = cur_latents - noise_pred_posi
|
||||
latents_total.append(cur_latents)
|
||||
LQ_pre_idx = LQ_cur_idx
|
||||
|
||||
|
||||
if unload_dit and hasattr(self, 'dit') and not next(self.dit.parameters()).is_cpu:
|
||||
try:
|
||||
del pre_cache_k, pre_cache_v
|
||||
except NameError:
|
||||
pass
|
||||
print("[FlashVSR] Offloading DiT to the CPU to free up VRAM...")
|
||||
self.dit.to('cpu')
|
||||
clean_vram()
|
||||
|
||||
latents = torch.cat(latents_total, dim=2)
|
||||
|
||||
del latents_total
|
||||
clean_vram()
|
||||
|
||||
if skip_vae:
|
||||
return latents
|
||||
|
||||
# Decode
|
||||
print("[FlashVSR] Starting VAE decoding...")
|
||||
frames = self.TCDecoder.decode_video(latents.transpose(1, 2),parallel=False, show_progress_bar=False, cond=LQ_video[:,:,:LQ_cur_idx,:,:]).transpose(1, 2).mul_(2).sub_(1)
|
||||
|
||||
# 颜色校正(wavelet)
|
||||
|
||||
@@ -18,6 +18,14 @@ from ..models.wan_video_vae import WanVideoVAE, RMS_norm, CausalConv3d, Upsample
|
||||
from ..schedulers.flow_match import FlowMatchScheduler
|
||||
from .base import BasePipeline
|
||||
|
||||
# --- Safetensors support ---
|
||||
st_load_file = None # Define the variable in global scope first
|
||||
try:
|
||||
from safetensors.torch import load_file as st_load_file
|
||||
except ImportError:
|
||||
# st_load_file remains None if import fails
|
||||
print("Warning: 'safetensors' not installed. Safetensors (.safetensors) files cannot be loaded.")
|
||||
# ---------------------------
|
||||
|
||||
# -----------------------------
|
||||
# 基础工具:ADAIN 所需的统计量(保留以备需要;管线默认用 wavelet)
|
||||
@@ -146,7 +154,7 @@ class TorchColorCorrectorWavelet(nn.Module):
|
||||
out = torch.cat(outs, dim=2)
|
||||
return out
|
||||
|
||||
|
||||
from ..models.utils import clean_vram
|
||||
# -----------------------------
|
||||
# 简化版 Pipeline(仅 dit + vae)
|
||||
# -----------------------------
|
||||
@@ -165,13 +173,12 @@ class FlashVSRTinyLongPipeline(BasePipeline):
|
||||
self.ColorCorrector = TorchColorCorrectorWavelet(levels=5)
|
||||
|
||||
print(r"""
|
||||
███████╗██╗ █████╗ ███████╗██╗ ██╗██╗ ██╗███████╗█████╗
|
||||
██╔════╝██║ ██╔══██╗██╔════╝██║ ██║██║ ██║██╔════╝██╔══██╗
|
||||
█████╗ ██║ ███████║███████╗███████║╚██╗ ██╔╝███████╗███████║
|
||||
██╔══╝ ██║ ██╔══██║╚════██║██╔══██║ ╚████╔╝ ╚════██║██╔═██║
|
||||
██║ ███████╗██║ ██║███████║██║ ██║ ╚██╔╝ ███████║██║ ██║
|
||||
╚═╝ ╚══════╝╚═╝ ╚═╝╚══════╝╚═╝ ╚═╝ ╚═╝ ╚══════╝╚═╝ ╚═╝
|
||||
⚡FlashVSR
|
||||
███████╗██╗ █████╗ ███████╗██╗ ██╗██╗ ██╗███████╗█████╗
|
||||
██╔════╝██║ ██╔══██╗██╔════╝██║ ██║██║ ██║██╔════╝██╔══██╗
|
||||
█████╗ ██║ ███████║███████╗███████║╚██╗ ██╔╝███████╗███████║
|
||||
██╔══╝ ██║ ██╔══██║╚════██║██╔══██║ ╚████╔╝ ╚════██║██╔═██║
|
||||
██║ ███████╗██║ ██║███████║██║ ██║ ╚██╔╝ ███████║██║ ██║
|
||||
╚═╝ ╚══════╝╚═╝ ╚═╝╚══════╝╚═╝ ╚═╝ ╚═╝ ╚══════╝╚═╝ ╚═╝
|
||||
""")
|
||||
|
||||
def enable_vram_management(self, num_persistent_param_in_dit=None):
|
||||
@@ -229,13 +236,14 @@ class FlashVSRTinyLongPipeline(BasePipeline):
|
||||
def init_cross_kv(
|
||||
self,
|
||||
context_tensor: Optional[torch.Tensor] = None,
|
||||
prompt_path = None,
|
||||
):
|
||||
self.load_models_to_device(["dit"])
|
||||
"""
|
||||
使用固定 prompt 生成文本 context,并在 WanModel 中初始化所有 CrossAttention 的 KV 缓存。
|
||||
必须在 __call__ 前显式调用一次。
|
||||
"""
|
||||
prompt_path = "../../examples/WanVSR/prompt_tensor/posi_prompt.pth"
|
||||
#prompt_path = "../../examples/WanVSR/prompt_tensor/posi_prompt.pth"
|
||||
|
||||
if self.dit is None:
|
||||
raise RuntimeError("请先通过 fetch_models / from_model_manager 初始化 self.dit")
|
||||
@@ -243,7 +251,30 @@ class FlashVSRTinyLongPipeline(BasePipeline):
|
||||
if context_tensor is None:
|
||||
if prompt_path is None:
|
||||
raise ValueError("init_cross_kv: 需要提供 prompt_path 或 context_tensor 其一")
|
||||
ctx = torch.load(prompt_path, map_location=self.device)
|
||||
|
||||
# --- Safetensors loading logic added here ---
|
||||
prompt_path_lower = prompt_path.lower()
|
||||
if prompt_path_lower.endswith(".safetensors"):
|
||||
if st_load_file is None:
|
||||
raise ImportError("The 'safetensors' library must be installed to load .safetensors files.")
|
||||
|
||||
# Load the tensor from safetensors
|
||||
loaded_dict = st_load_file(prompt_path, device=self.device)
|
||||
|
||||
# Safetensors loads a dict. Assuming the context tensor is the only or primary key.
|
||||
if len(loaded_dict) == 1:
|
||||
ctx = list(loaded_dict.values())[0]
|
||||
elif 'context' in loaded_dict: # Common key for text context
|
||||
ctx = loaded_dict['context']
|
||||
else:
|
||||
raise ValueError(f"Safetensors file {prompt_path} does not contain an obvious single tensor ('context' key not found and multiple keys exist).")
|
||||
|
||||
else:
|
||||
# Default behavior for .pth, .pt, etc.
|
||||
ctx = torch.load(prompt_path, map_location=self.device)
|
||||
|
||||
# --------------------------------------------
|
||||
# ctx = torch.load(prompt_path, map_location=self.device)
|
||||
else:
|
||||
ctx = context_tensor
|
||||
|
||||
@@ -274,9 +305,19 @@ class FlashVSRTinyLongPipeline(BasePipeline):
|
||||
latents = self.vae.encode(input_video, device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
|
||||
return latents
|
||||
|
||||
def decode_video(self, latents, tiled=True, tile_size=(34, 34), tile_stride=(18, 16)):
|
||||
def _decode_video(self, latents, tiled=True, tile_size=(34, 34), tile_stride=(18, 16)):
|
||||
frames = self.vae.decode(latents, device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
|
||||
return frames
|
||||
|
||||
def decode_video(self, latents, cond=None, **kwargs):
|
||||
frames = self.TCDecoder.decode_video(
|
||||
latents.transpose(1, 2), # TCDecoder 需要 (B, F, C, H, W)
|
||||
parallel=False,
|
||||
show_progress_bar=False,
|
||||
cond=cond
|
||||
).transpose(1, 2).mul_(2).sub_(1) # 转回 (B, C, F, H, W) 格式,范围 -1 to 1
|
||||
|
||||
return frames
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
@@ -296,7 +337,7 @@ class FlashVSRTinyLongPipeline(BasePipeline):
|
||||
tile_size=(60, 104),
|
||||
tile_stride=(30, 52),
|
||||
tea_cache_l1_thresh=None,
|
||||
tea_cache_model_id="Wan2.1-T2V-14B",
|
||||
tea_cache_model_id="Wan2.1-T2V-1.3B",
|
||||
progress_bar_cmd=tqdm,
|
||||
progress_bar_st=None,
|
||||
LQ_video=None,
|
||||
@@ -306,10 +347,12 @@ class FlashVSRTinyLongPipeline(BasePipeline):
|
||||
kv_ratio=3.0,
|
||||
local_range = 9,
|
||||
color_fix = True,
|
||||
unload_dit = False,
|
||||
skip_vae = False,
|
||||
):
|
||||
# 只接受 cfg=1.0(与原代码一致)
|
||||
assert cfg_scale == 1.0, "cfg_scale must be 1.0"
|
||||
|
||||
|
||||
# 要求:必须先 init_cross_kv()
|
||||
if self.prompt_emb_posi is None or 'context' not in self.prompt_emb_posi:
|
||||
raise RuntimeError(
|
||||
@@ -343,13 +386,13 @@ class FlashVSRTinyLongPipeline(BasePipeline):
|
||||
if hasattr(self.dit, "LQ_proj_in"):
|
||||
self.dit.LQ_proj_in.clear_cache()
|
||||
|
||||
self.TCDecoder.clean_mem()
|
||||
frames_total = []
|
||||
LQ_pre_idx = 0
|
||||
LQ_cur_idx = 0
|
||||
frames_total = []
|
||||
self.TCDecoder.clean_mem()
|
||||
|
||||
with torch.no_grad():
|
||||
for cur_process_idx in tqdm(range(process_total_num)):
|
||||
for cur_process_idx in progress_bar_cmd(range(process_total_num)):
|
||||
if cur_process_idx == 0:
|
||||
pre_cache_k = [None] * len(self.dit.blocks)
|
||||
pre_cache_v = [None] * len(self.dit.blocks)
|
||||
@@ -384,7 +427,7 @@ class FlashVSRTinyLongPipeline(BasePipeline):
|
||||
LQ_latents[layer_idx] = torch.cat([LQ_latents[layer_idx], cur[layer_idx]], dim=1)
|
||||
LQ_cur_idx = cur_process_idx*8+21+(inner_loop_num-2)*4
|
||||
cur_latents = latents[:, :, 4+cur_process_idx*2:6+cur_process_idx*2, :, :]
|
||||
|
||||
|
||||
# 推理(无 motion_controller / vace)
|
||||
noise_pred_posi, pre_cache_k, pre_cache_v = model_fn_wan_video(
|
||||
self.dit,
|
||||
@@ -408,9 +451,14 @@ class FlashVSRTinyLongPipeline(BasePipeline):
|
||||
|
||||
# 更新 latent
|
||||
cur_latents = cur_latents - noise_pred_posi
|
||||
|
||||
# Decode
|
||||
cur_LQ_frame = LQ_video[:,:,LQ_pre_idx:LQ_cur_idx,:,:].to(self.device)
|
||||
cur_frames = self.TCDecoder.decode_video(cur_latents.transpose(1, 2),parallel=False, show_progress_bar=False, cond=LQ_video[:,:,LQ_pre_idx:LQ_cur_idx,:,:].to(self.device)).transpose(1, 2).mul_(2).sub_(1)
|
||||
cur_frames = self.TCDecoder.decode_video(
|
||||
cur_latents.transpose(1, 2),
|
||||
parallel=False,
|
||||
show_progress_bar=False,
|
||||
cond=cur_LQ_frame).transpose(1, 2).mul_(2).sub_(1)
|
||||
|
||||
# 颜色校正(wavelet)
|
||||
try:
|
||||
@@ -424,12 +472,14 @@ class FlashVSRTinyLongPipeline(BasePipeline):
|
||||
)
|
||||
except:
|
||||
pass
|
||||
|
||||
|
||||
frames_total.append(cur_frames.to('cpu'))
|
||||
LQ_pre_idx = LQ_cur_idx
|
||||
|
||||
|
||||
del cur_frames, cur_latents, cur_LQ_frame
|
||||
clean_vram()
|
||||
|
||||
frames = torch.cat(frames_total, dim=2)
|
||||
|
||||
return frames[0]
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user