@@ -25,12 +25,173 @@ from __future__ import annotations
|
||||
import functools
|
||||
import itertools
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from typing import Any,TypeVar
|
||||
import gc
|
||||
import torch
|
||||
from torch import nn
|
||||
from contextlib import contextmanager
|
||||
from collections.abc import Iterator
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
_M = TypeVar("_M", bound=torch.nn.Module)
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def cleanup_memory() -> None:
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# LayerStreamingWrapper from https://github.com/Lightricks/LTX-2
|
||||
|
||||
class SimpleLayerStreamingWrapper_Dual(nn.Module):
|
||||
"""简化版层流式处理包装器,支持多模块卸载"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: nn.Module,
|
||||
layers_attrs: list[str], # 修改为列表,支持多个模块路径
|
||||
target_device: torch.device,
|
||||
active_count: int = 1,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self._model = model
|
||||
self._layers_attrs = layers_attrs
|
||||
self._target_device = target_device
|
||||
self._active_count = active_count
|
||||
|
||||
# 解析并存储所有需要卸载的模块
|
||||
self._layer_groups: list[nn.ModuleList] = []
|
||||
self._stores: list[_SimpleLayerStore] = []
|
||||
|
||||
for attr in self._layers_attrs:
|
||||
layers = _resolve_attr(model, attr)
|
||||
self._layer_groups.append(layers)
|
||||
self._stores.append(_SimpleLayerStore(layers, self._target_device))
|
||||
|
||||
# 将非层参数移到GPU
|
||||
self._move_non_layer_params_to_gpu()
|
||||
|
||||
# 为所有模块组注册钩子
|
||||
self._register_simple_hooks()
|
||||
|
||||
def _move_non_layer_params_to_gpu(self) -> None:
|
||||
"""移动非层参数到GPU,排除所有需要流式卸载的模块参数"""
|
||||
layer_tensor_ids = set()
|
||||
# 收集所有卸载模块的参数 ID
|
||||
for layers in self._layer_groups:
|
||||
for layer in layers:
|
||||
for t in itertools.chain(layer.parameters(), layer.buffers()):
|
||||
layer_tensor_ids.add(id(t))
|
||||
|
||||
for p in self._model.parameters():
|
||||
if id(p) not in layer_tensor_ids:
|
||||
p.data = p.data.to(self._target_device)
|
||||
for b in self._model.buffers():
|
||||
if id(b) not in layer_tensor_ids:
|
||||
b.data = b.data.to(self._target_device)
|
||||
def forward(self, *args: Any, **kwargs: Any) -> Any:
|
||||
return self._model(*args, **kwargs)
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
"""代理属性访问到原始模型"""
|
||||
try:
|
||||
# 首先尝试从包装器自身获取属性
|
||||
return super().__getattr__(name)
|
||||
except AttributeError:
|
||||
# 如果失败,则从原始模型获取
|
||||
return getattr(self._model, name)
|
||||
|
||||
def _register_simple_hooks(self) -> None:
|
||||
"""为所有模块组注册简单的加载/释放钩子"""
|
||||
# 遍历每一个模块组及其对应的 Store
|
||||
for layers, store in zip(self._layer_groups, self._stores):
|
||||
idx_map = {id(layer): idx for idx, layer in enumerate(layers)}
|
||||
|
||||
def _pre_hook(module: nn.Module, input, *, idx: int, s: _SimpleLayerStore):
|
||||
# 加载当前层到GPU
|
||||
s.load_layer_to_gpu(idx, module)
|
||||
# 记录流,防止内存被提前回收
|
||||
for param in itertools.chain(module.parameters(), module.buffers()):
|
||||
param.data.record_stream(torch.cuda.current_stream(self._target_device))
|
||||
|
||||
def _post_hook(module: nn.Module, input, output, *, idx: int, s: _SimpleLayerStore):
|
||||
# 处理完后立即将层移回CPU
|
||||
s.unload_layer_from_gpu(idx, module)
|
||||
|
||||
for layer in layers:
|
||||
idx = idx_map[id(layer)]
|
||||
# 使用 functools.partial 将对应的 store 实例传入钩子
|
||||
pre_hook = layer.register_forward_pre_hook(functools.partial(_pre_hook, idx=idx, s=store))
|
||||
post_hook = layer.register_forward_hook(functools.partial(_post_hook, idx=idx, s=store))
|
||||
|
||||
@contextmanager
|
||||
def _streaming_model(
|
||||
model: _M,
|
||||
layers_attr, # 允许接收 str 或 list[str]
|
||||
target_device: torch.device,
|
||||
prefetch_count: int,
|
||||
) -> Iterator[_M]:
|
||||
"""Wrap *model* with :class:`LayerStreamingWrapper`, yield it, then tear down."""
|
||||
# 根据传入的 layers_attr 类型自动路由到对应的 Wrapper
|
||||
if isinstance(layers_attr, list):
|
||||
wrapped = SimpleLayerStreamingWrapper_Dual(
|
||||
model,
|
||||
layers_attrs=layers_attr,
|
||||
target_device=target_device,
|
||||
active_count=prefetch_count,
|
||||
)
|
||||
else:
|
||||
wrapped = SimpleLayerStreamingWrapper(
|
||||
model,
|
||||
layers_attr=layers_attr,
|
||||
target_device=target_device,
|
||||
active_count=prefetch_count,
|
||||
)
|
||||
|
||||
try:
|
||||
yield wrapped # type: ignore[misc]
|
||||
finally:
|
||||
wrapped.to("cpu")
|
||||
cleanup_memory()
|
||||
torch.cuda.synchronize(device=target_device)
|
||||
try:
|
||||
if hasattr(torch._C, "_host_emptyCache"):
|
||||
torch._C._host_emptyCache()
|
||||
except Exception:
|
||||
print("Host empty cache cleanup failed; ignoring.", exc_info=True)
|
||||
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _streaming_model_(
|
||||
model: _M,
|
||||
layers_attr: str,
|
||||
target_device: torch.device,
|
||||
prefetch_count: int,
|
||||
) -> Iterator[_M]:
|
||||
"""Wrap *model* with :class:`LayerStreamingWrapper`, yield it, then tear down."""
|
||||
wrapped = SimpleLayerStreamingWrapper(
|
||||
model,
|
||||
layers_attr=layers_attr,
|
||||
target_device=target_device,
|
||||
active_count=prefetch_count,
|
||||
)
|
||||
try:
|
||||
yield wrapped # type: ignore[misc]
|
||||
finally:
|
||||
wrapped.to("cpu")
|
||||
cleanup_memory()
|
||||
# Flush the host (pinned) memory cache so that freed pinned pages are
|
||||
# returned to the OS. Without this, sequential streaming models
|
||||
# (e.g. text encoder then transformer) exhaust host memory because the
|
||||
# CachingHostAllocator keeps freed blocks cached indefinitely.
|
||||
torch.cuda.synchronize(device=target_device)
|
||||
try:
|
||||
if hasattr(torch._C, "_host_emptyCache"):
|
||||
torch._C._host_emptyCache()
|
||||
except Exception:
|
||||
print("Host empty cache cleanup failed; ignoring.", exc_info=True)
|
||||
|
||||
|
||||
def _resolve_attr(module: nn.Module, dotted_path: str) -> nn.ModuleList:
|
||||
|
||||
@@ -329,6 +329,39 @@ class LongCatVideoAvatarTransformer3DModel(
|
||||
module.forward = self._create_multi_lora_forward(module, loras)
|
||||
|
||||
def _create_multi_lora_forward(self, module, loras):
|
||||
def multi_lora_forward(x, *args, **kwargs):
|
||||
weight_dtype = x.dtype
|
||||
target_device = x.device
|
||||
|
||||
# 执行原始的模块前向传播(不受LoRA设备影响)
|
||||
org_output = module.org_forward(x, *args, **kwargs)
|
||||
|
||||
total_lora_output = 0
|
||||
for lora in loras:
|
||||
if lora.use_lora:
|
||||
# 1. 推理前:将 LoRA 权重动态加载到输入张量所在的 CUDA 设备
|
||||
lora.lora_down.to(target_device, dtype=weight_dtype, non_blocking=True)
|
||||
lora.lora_up.to(target_device, dtype=weight_dtype, non_blocking=True)
|
||||
|
||||
# 2. 执行 LoRA 计算
|
||||
lx = lora.lora_down(x)
|
||||
lx = lora.lora_up(lx)
|
||||
lora_output = lx * lora.multiplier * lora.alpha_scale
|
||||
|
||||
total_lora_output += lora_output
|
||||
|
||||
# 3. 推理后:立即将 LoRA 权重卸载回 CPU,释放显存
|
||||
lora.lora_down.to("cpu", non_blocking=True)
|
||||
lora.lora_up.to("cpu", non_blocking=True)
|
||||
|
||||
# 累加 LoRA 输出并转换回原始数据类型
|
||||
total_lora_output = total_lora_output.to(weight_dtype)
|
||||
|
||||
return org_output + total_lora_output
|
||||
|
||||
return multi_lora_forward
|
||||
|
||||
def _create_multi_lora_forward_(self, module, loras):
|
||||
def multi_lora_forward(x, *args, **kwargs):
|
||||
weight_dtype = x.dtype
|
||||
org_output = module.org_forward(x, *args, **kwargs)
|
||||
|
||||
@@ -6,7 +6,7 @@ from typing import Optional, Set
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from safetensors.torch import save_file, load_file
|
||||
from safetensors.torch import save_file, load_file as safe_load_file
|
||||
|
||||
|
||||
class QuantizedLinear(nn.Module):
|
||||
@@ -233,17 +233,17 @@ def load_quantized_dit(checkpoint_dir: str, subfolder: str = "base_model_int8",
|
||||
state_dict = {}
|
||||
for shard_file in sorted(shard_files):
|
||||
shard_path = os.path.join(quantized_dir, shard_file)
|
||||
shard_dict = load_file(shard_path, device="cpu")
|
||||
shard_dict = safe_load_file(shard_path, device="cpu")
|
||||
state_dict.update(shard_dict)
|
||||
else:
|
||||
# Single file fallback
|
||||
if single_file:
|
||||
state_dict = load_file(single_file, device="cpu")
|
||||
state_dict = safe_load_file(single_file, device="cpu")
|
||||
else:
|
||||
files = [f for f in os.listdir(quantized_dir) if f.endswith(".safetensors") and "index" not in f]
|
||||
state_dict = {}
|
||||
for f in sorted(files):
|
||||
shard_dict = load_file(os.path.join(quantized_dir, f), device="cpu")
|
||||
shard_dict = safe_load_file(os.path.join(quantized_dir, f), device="cpu")
|
||||
state_dict.update(shard_dict)
|
||||
|
||||
X=model.load_state_dict(state_dict, strict=True,assign=True) # Load weights and cast to bfloat16 for non-quantized params
|
||||
|
||||
@@ -19,7 +19,7 @@ from .modules.autoencoder_kl_wan import AutoencoderKLWan
|
||||
from .modules.avatar.longcat_video_dit_avatar import LongCatVideoAvatarTransformer3DModel
|
||||
#from .context_parallel import context_parallel_util
|
||||
from .utils.bukcet_config import get_bucket_config
|
||||
from ..utils import _streaming_model
|
||||
from ..layer_streaming import _streaming_model
|
||||
from contextlib import AbstractContextManager
|
||||
import ftfy
|
||||
import regex as re
|
||||
@@ -36,6 +36,67 @@ def torch_gc():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
@torch.no_grad()
|
||||
def get_audio_embedding_whisper_(audio_encoder,speech_array, fps=25, device='cpu', sample_rate=16000):
|
||||
"""使用 Whisper encoder 提取音频特征。
|
||||
Args:
|
||||
speech_array: 原始音频波形 (numpy array, 单声道, sample_rate=16000)
|
||||
fps: 目标视频帧率
|
||||
device: 推理设备
|
||||
sample_rate: 音频采样率
|
||||
Returns:
|
||||
audio_emb: [T, 5, D],T = int(audio_duration * fps)
|
||||
"""
|
||||
def linear_interpolation_fps(features, input_fps, output_fps, output_len=None):
|
||||
"""将音频特征从 input_fps 插值到 output_fps。
|
||||
Args:
|
||||
features: [B, T, D]
|
||||
input_fps: 源帧率
|
||||
output_fps: 目标帧率
|
||||
output_len: 若指定则直接用该长度,否则按帧率比换算
|
||||
Returns:
|
||||
[B, output_len, D]
|
||||
"""
|
||||
features = features.transpose(1, 2) # [B, D, T]
|
||||
if output_len is None:
|
||||
output_len = int(features.shape[2] / float(input_fps) * output_fps)
|
||||
output_features = F.interpolate(features, size=output_len, align_corners=True, mode='linear')
|
||||
return output_features.transpose(1, 2)
|
||||
# ---- 常量 ----
|
||||
MEL_CHUNK = 750 * 640 # feature extractor 滑窗大小(样本数)
|
||||
ENC_CHUNK = 3000 # encoder 滑窗大小(mel 帧数)
|
||||
ENC_FPS = 50 # encoder 输出帧率
|
||||
|
||||
|
||||
audio_duration = len(speech_array) / sample_rate
|
||||
video_length = int(audio_duration * fps)
|
||||
def _loudness_norm(audio_array, sr=16000, lufs=-23, threshold=100):
|
||||
meter = pyln.Meter(sr)
|
||||
loudness = meter.integrated_loudness(audio_array)
|
||||
if abs(loudness) > threshold:
|
||||
return audio_array
|
||||
normalized_audio = pyln.normalize.loudness(audio_array, loudness, lufs)
|
||||
return normalized_audio
|
||||
# ---- 音频预处理 ----
|
||||
speech_array = _loudness_norm(speech_array, sample_rate)
|
||||
|
||||
|
||||
# comfyui encoder
|
||||
audio_encoder_output = audio_encoder.encode_audio(torch.from_numpy(speech_array).unsqueeze(0).unsqueeze(0).to(device,audio_encoder.dtype), sample_rate)
|
||||
audio_emb = torch.stack(audio_encoder_output["encoded_audio_all_layers"], dim=2)
|
||||
audio_len = audio_encoder_output["audio_samples"] // 640
|
||||
|
||||
audio_prompts = audio_emb[:, :audio_len * 2]
|
||||
|
||||
# ---- 按层分组 + 插值到目标帧数 ----
|
||||
feat0 = linear_interpolation_fps(audio_prompts[:, :, 0: 8].mean(dim=2), ENC_FPS, fps, video_length)
|
||||
feat1 = linear_interpolation_fps(audio_prompts[:, :, 8:16].mean(dim=2), ENC_FPS, fps, video_length)
|
||||
feat2 = linear_interpolation_fps(audio_prompts[:, :, 16:24].mean(dim=2), ENC_FPS, fps, video_length)
|
||||
feat3 = linear_interpolation_fps(audio_prompts[:, :, 24:32].mean(dim=2), ENC_FPS, fps, video_length)
|
||||
feat4 = linear_interpolation_fps(audio_prompts[:, :, 32], ENC_FPS, fps, video_length)
|
||||
audio_emb = torch.stack([feat0, feat1, feat2, feat3, feat4], dim=2)[0] # [T, 5, D]
|
||||
return audio_emb
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def get_audio_embedding_whisper(audio_encoder,audio_feature_extractor,speech_array, fps=25, device='cpu', sample_rate=16000):
|
||||
@@ -80,7 +141,7 @@ def get_audio_embedding_whisper(audio_encoder,audio_feature_extractor,speech_arr
|
||||
return normalized_audio
|
||||
# ---- 音频预处理 ----
|
||||
speech_array = _loudness_norm(speech_array, sample_rate)
|
||||
|
||||
|
||||
# ---- Whisper feature extractor:wav → mel spectrogram ----
|
||||
mel_chunks = []
|
||||
for i in range(0, len(speech_array), MEL_CHUNK):
|
||||
@@ -101,8 +162,11 @@ def get_audio_embedding_whisper(audio_encoder,audio_feature_extractor,speech_arr
|
||||
output_hidden_states=True,
|
||||
).hidden_states # tuple: (n_layers+1,) x [1, T_enc, D]
|
||||
enc_chunks.append(torch.stack(chunk_hs, dim=2)) # [1, T_enc, n_layers, D]
|
||||
|
||||
audio_prompts = torch.cat(enc_chunks, dim=1) # [1, T_enc_total, n_layers, D]
|
||||
|
||||
audio_prompts = audio_prompts[:, :video_length * 2] # 截取有效帧
|
||||
|
||||
|
||||
# ---- 按层分组 + 插值到目标帧数 ----
|
||||
feat0 = linear_interpolation_fps(audio_prompts[:, :, 0: 8].mean(dim=2), ENC_FPS, fps, video_length)
|
||||
@@ -250,7 +314,7 @@ class LongCatVideoAvatarPipeline:
|
||||
if streaming_prefetch_count is not None:
|
||||
return _streaming_model(
|
||||
self.dit,
|
||||
layers_attr="blocks",
|
||||
layers_attr=["blocks"],
|
||||
target_device=torch.device("cuda"),
|
||||
prefetch_count=streaming_prefetch_count,
|
||||
)
|
||||
@@ -1665,10 +1729,10 @@ class LongCatVideoAvatarPipeline:
|
||||
self.device = device
|
||||
if self.vae is not None:
|
||||
self.vae = self.vae.to(device, non_blocking=True)
|
||||
if hasattr(self.dit, 'lora_dict') and self.dit.lora_dict:
|
||||
for lora_key, lora_network in self.dit.lora_dict.items():
|
||||
for lora in lora_network.loras:
|
||||
lora.to(device, non_blocking=True)
|
||||
# if hasattr(self.dit, 'lora_dict') and self.dit.lora_dict:
|
||||
# for lora_key, lora_network in self.dit.lora_dict.items():
|
||||
# for lora in lora_network.loras:
|
||||
# lora.to(device, non_blocking=True)
|
||||
return self
|
||||
|
||||
def to(self, device: str | torch.device):
|
||||
|
||||
@@ -426,7 +426,7 @@ def generate_multi(pipe,condition,te_cond,device,seed,cond_image,resolution,
|
||||
generator=generator,
|
||||
output_type='both',
|
||||
use_kv_cache=True,
|
||||
offload_kv_cache=False,
|
||||
offload_kv_cache=True,
|
||||
enhance_hf=True if not use_distill else False,
|
||||
audio_emb=audio_embs,
|
||||
ref_latent=ref_latent,
|
||||
|
||||
@@ -15,7 +15,7 @@ import torch
|
||||
from transformers import AutoTokenizer, UMT5EncoderModel
|
||||
from diffusers.utils import load_image
|
||||
|
||||
from .longcat_video.pipeline_longcat_video_avatar import LongCatVideoAvatarPipeline,get_audio_embedding_whisper
|
||||
from .longcat_video.pipeline_longcat_video_avatar import LongCatVideoAvatarPipeline,get_audio_embedding_whisper,get_audio_embedding_whisper_
|
||||
from .longcat_video.modules.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
|
||||
from .longcat_video.modules.autoencoder_kl_wan import AutoencoderKLWan
|
||||
from .longcat_video.modules.avatar.longcat_video_dit_avatar import LongCatVideoAvatarTransformer3DModel
|
||||
@@ -47,7 +47,7 @@ def extract_vocal_from_speech(source_path, target_path, vocal_separator, audio_o
|
||||
print("Audio separate failed. Using raw audio.")
|
||||
return None
|
||||
|
||||
default_vocal_path = audio_output_dir_temp / "vocals" / outputs[0]
|
||||
default_vocal_path = Path(os.path.join(audio_output_dir_temp ,"vocals", f"{outputs[0]}"))
|
||||
default_vocal_path = default_vocal_path.resolve().as_posix()
|
||||
# cmd = f"mv '{default_vocal_path}' '{target_path}'"
|
||||
# os.system(cmd)
|
||||
@@ -178,7 +178,7 @@ def prepare_audio(audio, sample_rate=16000):
|
||||
sr = 16000
|
||||
return speech_array,sr
|
||||
|
||||
def get_audio_emb(checkpoint_dir,audio,left_audio,audio_type,save_fps,num_segments,device,p_box,model_type='avatar-v1.5' ):
|
||||
def get_audio_emb(audio_encoder,audio,left_audio,audio_type,save_fps,num_segments,device,p_box,model_type='avatar-v1.5' ):
|
||||
num_frames=93
|
||||
num_cond_frames = 13
|
||||
audio_stride=1
|
||||
@@ -198,23 +198,16 @@ def get_audio_emb(checkpoint_dir,audio,left_audio,audio_type,save_fps,num_segmen
|
||||
left_person_bbox=p_box[0]
|
||||
right_person_bbox=p_box[1]
|
||||
other_person_bbox = p_box[2] if len(p_box) > 2 else None
|
||||
# left_person_bbox = p_box.get('person1', None)
|
||||
# right_person_bbox = p_box.get('person2', None)
|
||||
# other_person_bbox = p_box.get('others', None)
|
||||
use_background_silent_audio = other_person_bbox is not None and len(other_person_bbox) > 0
|
||||
# audio embedding
|
||||
# initialize audio models
|
||||
audio_model_checkpoint_path = os.path.join(checkpoint_dir, 'whisper-large-v3')
|
||||
audio_encoder = get_audio_encoder(audio_model_checkpoint_path, model_type).to(device)
|
||||
audio_feature_extractor = get_audio_feature_extractor(audio_model_checkpoint_path, model_type)
|
||||
|
||||
if left_speech_array is not None:
|
||||
left_speech_array_ext, right_speech_array_ext = audio_prepare_multi(left_speech_array,speech_array, generate_duration, sr=sr, audio_type=audio_type)
|
||||
left_full_audio_emb = get_audio_embedding_whisper(audio_encoder, audio_feature_extractor, left_speech_array_ext, fps=save_fps*audio_stride, device="cuda" if torch.cuda.is_available() else "cpu", sample_rate=sr)
|
||||
full_audio_emb = get_audio_embedding_whisper(audio_encoder, audio_feature_extractor, right_speech_array_ext, fps=save_fps*audio_stride, device="cuda" if torch.cuda.is_available() else "cpu", sample_rate=sr)
|
||||
left_full_audio_emb = get_audio_embedding_whisper_(audio_encoder, left_speech_array_ext, fps=save_fps*audio_stride, )
|
||||
full_audio_emb = get_audio_embedding_whisper_(audio_encoder, right_speech_array_ext, fps=save_fps*audio_stride,)
|
||||
if torch.isnan(left_full_audio_emb).any() or torch.isnan(full_audio_emb).any():
|
||||
raise ValueError(f"broken audio embedding with nan values")
|
||||
if use_background_silent_audio:
|
||||
back_full_audio_emb = get_audio_embedding_whisper(audio_encoder, audio_feature_extractor,np.zeros_like(left_speech_array_ext), fps=save_fps*audio_stride, device="cuda" if torch.cuda.is_available() else "cpu", sample_rate=sr)
|
||||
back_full_audio_emb = get_audio_embedding_whisper_(audio_encoder,np.zeros_like(left_speech_array_ext), fps=save_fps*audio_stride, )
|
||||
assert left_full_audio_emb.shape == full_audio_emb.shape, f"Inconsistent audio embedding shape."
|
||||
if use_background_silent_audio:
|
||||
assert left_full_audio_emb.shape == back_full_audio_emb.shape, f"Inconsistent audio embedding shape between speaker and background."
|
||||
@@ -223,19 +216,10 @@ def get_audio_emb(checkpoint_dir,audio,left_audio,audio_type,save_fps,num_segmen
|
||||
added_sample_nums = math.ceil((generate_duration - source_duraion) * sr)
|
||||
if added_sample_nums > 0:
|
||||
speech_array = np.append(speech_array, [0.]*added_sample_nums)
|
||||
full_audio_emb = get_audio_embedding_whisper(audio_encoder, audio_feature_extractor, speech_array, fps=save_fps*audio_stride, device="cuda" if torch.cuda.is_available() else "cpu", sample_rate=sr)
|
||||
|
||||
full_audio_emb=get_audio_embedding_whisper_(audio_encoder, speech_array, fps=save_fps*audio_stride, ) #torch.Size([2142, 5, 1280])
|
||||
if torch.isnan(full_audio_emb).any():
|
||||
raise ValueError(f"broken audio embedding with nan values")
|
||||
|
||||
# # prepare audio embedding for the first clip
|
||||
# indices = torch.arange(2 * 2 + 1) - 2
|
||||
# audio_start_idx = 0
|
||||
# audio_end_idx = audio_start_idx + audio_stride * num_frames
|
||||
|
||||
# center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
|
||||
# center_indices = torch.clamp(center_indices, min=0, max=full_audio_emb.shape[0]-1)
|
||||
# audio_emb = full_audio_emb[center_indices][None,...].to(device)
|
||||
au_cond={
|
||||
"full_audio_emb": full_audio_emb,
|
||||
"num_segments": num_segments,
|
||||
@@ -250,8 +234,9 @@ def get_audio_emb(checkpoint_dir,audio,left_audio,audio_type,save_fps,num_segmen
|
||||
return au_cond
|
||||
|
||||
|
||||
def get_audio_vocal(checkpoint_dir,raw_speech_path,left_raw_audio_path,audio_output_dir_temp,):
|
||||
vocal_separator_path = os.path.join(checkpoint_dir, 'vocal_separator/Kim_Vocal_2.onnx')
|
||||
def load_audio_vocal(vocal_separator_path,audio_output_dir_temp,checkpoint_dir):
|
||||
if vocal_separator_path is None:
|
||||
vocal_separator_path = os.path.join(checkpoint_dir, 'Kim_Vocal_2.onnx')
|
||||
os.makedirs(audio_output_dir_temp, exist_ok=True)
|
||||
audio_output_dir_temp = Path(audio_output_dir_temp)
|
||||
audio_separator_model_path = os.path.dirname(vocal_separator_path)
|
||||
@@ -264,33 +249,23 @@ def get_audio_vocal(checkpoint_dir,raw_speech_path,left_raw_audio_path,audio_out
|
||||
|
||||
vocal_separator.load_model(audio_separator_model_name)
|
||||
vocal_separator.onnx_execution_provider = ["CUDAExecutionProvider"]
|
||||
return vocal_separator
|
||||
|
||||
|
||||
def get_audio_vocal(vocal_separator,raw_speech_path,audio_output_dir_temp,):
|
||||
|
||||
vocal_path=replace_to_vocal_suffix(raw_speech_path)
|
||||
left_audio_path=replace_to_vocal_suffix(left_raw_audio_path) if left_raw_audio_path is not None else None
|
||||
os.makedirs(os.path.dirname(vocal_path), exist_ok=True)
|
||||
|
||||
temp_vocal_path = extract_vocal_from_speech(raw_speech_path,vocal_path , vocal_separator, audio_output_dir_temp)
|
||||
|
||||
temp_left_vocal_path,left_audio=None,None
|
||||
|
||||
if left_audio_path is not None:
|
||||
os.makedirs(os.path.dirname(temp_vocal_path), exist_ok=True)
|
||||
temp_left_vocal_path = extract_vocal_from_speech(left_raw_audio_path,left_audio_path , vocal_separator, audio_output_dir_temp)
|
||||
|
||||
temp_vocal_path = extract_vocal_from_speech(raw_speech_path,vocal_path , vocal_separator, audio_output_dir_temp)
|
||||
import librosa
|
||||
vocal_array, sr = librosa.load(temp_vocal_path, sr=16000)
|
||||
if temp_left_vocal_path is not None:
|
||||
left_vocal_array, sr = librosa.load(temp_left_vocal_path, sr=16000)
|
||||
left_audio={
|
||||
"waveform": torch.from_numpy(left_vocal_array).unsqueeze(0).unsqueeze(0),
|
||||
"sample_rate": sr,
|
||||
}
|
||||
#print("vocal_array.shape", vocal_array.shape)
|
||||
audio={
|
||||
"waveform": torch.from_numpy(vocal_array).unsqueeze(0).unsqueeze(0),
|
||||
"sample_rate": sr,
|
||||
}
|
||||
return temp_vocal_path,audio,left_audio
|
||||
return temp_vocal_path,audio
|
||||
|
||||
def generate(pipe,condition,te_cond,device,seed,stage_1,cond_image,resolution,
|
||||
text_guidance_scale,audio_guidance_scale,num_inference_steps,ref_img_index,mask_frame_range,
|
||||
@@ -434,7 +409,7 @@ def generate(pipe,condition,te_cond,device,seed,stage_1,cond_image,resolution,
|
||||
generator=generator,
|
||||
output_type='both',
|
||||
use_kv_cache=True,
|
||||
offload_kv_cache=False,
|
||||
offload_kv_cache=True,
|
||||
enhance_hf=True if not use_distill else False,
|
||||
audio_emb=audio_emb,
|
||||
ref_latent=ref_latent,
|
||||
|
||||
+1
-45
@@ -1,57 +1,12 @@
|
||||
from diffusers.quantizers.gguf.utils import dequantize_gguf_tensor
|
||||
|
||||
from contextlib import contextmanager
|
||||
from .layer_streaming import SimpleLayerStreamingWrapper
|
||||
from collections.abc import Iterator
|
||||
from typing import TypeVar
|
||||
import gc
|
||||
import torch
|
||||
# from utils import apply_loras_gguf
|
||||
|
||||
_M = TypeVar("_M", bound=torch.nn.Module)
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
|
||||
def cleanup_memory() -> None:
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# LayerStreamingWrapper from https://github.com/Lightricks/LTX-2
|
||||
|
||||
@contextmanager
|
||||
def _streaming_model(
|
||||
model: _M,
|
||||
layers_attr: str,
|
||||
target_device: torch.device,
|
||||
prefetch_count: int,
|
||||
) -> Iterator[_M]:
|
||||
"""Wrap *model* with :class:`LayerStreamingWrapper`, yield it, then tear down."""
|
||||
wrapped = SimpleLayerStreamingWrapper(
|
||||
model,
|
||||
layers_attr=layers_attr,
|
||||
target_device=target_device,
|
||||
active_count=prefetch_count,
|
||||
)
|
||||
try:
|
||||
yield wrapped # type: ignore[misc]
|
||||
finally:
|
||||
wrapped.to("cpu")
|
||||
cleanup_memory()
|
||||
# Flush the host (pinned) memory cache so that freed pinned pages are
|
||||
# returned to the OS. Without this, sequential streaming models
|
||||
# (e.g. text encoder then transformer) exhaust host memory because the
|
||||
# CachingHostAllocator keeps freed blocks cached indefinitely.
|
||||
torch.cuda.synchronize(device=target_device)
|
||||
try:
|
||||
if hasattr(torch._C, "_host_emptyCache"):
|
||||
torch._C._host_emptyCache()
|
||||
except Exception:
|
||||
print("Host empty cache cleanup failed; ignoring.", exc_info=True)
|
||||
|
||||
|
||||
|
||||
def set_gguf2meta_model(meta_model,model_state_dict,dtype,device,lora_sd=None):
|
||||
from diffusers import GGUFQuantizationConfig
|
||||
from diffusers.quantizers.gguf import GGUFQuantizer
|
||||
@@ -153,6 +108,7 @@ def apply_loras_gguf(
|
||||
model_sd,
|
||||
lora_sd,
|
||||
):
|
||||
from diffusers.quantizers.gguf.utils import dequantize_gguf_tensor
|
||||
sd = {}
|
||||
for key, weight in model_sd.items():
|
||||
if weight is None:
|
||||
|
||||
+30
-9
@@ -8,7 +8,7 @@ import os
|
||||
from comfy_api.latest import io
|
||||
import folder_paths
|
||||
from .node_utils import clear_comfyui_cache,tensor2image,audio2path
|
||||
from .LongCat_Video.run_demo_avatar_single_audio_to_video import load_longcat_video_model,generate,get_audio_vocal,get_audio_emb
|
||||
from .LongCat_Video.run_demo_avatar_single_audio_to_video import load_longcat_video_model,generate,get_audio_vocal,get_audio_emb,load_audio_vocal
|
||||
from .LongCat_Video.run_demo_avatar_multi_audio_to_video import generate_multi
|
||||
device = torch.device(
|
||||
"cuda:0") if torch.cuda.is_available() else torch.device(
|
||||
@@ -145,6 +145,7 @@ class LongCat_Video_SM_Audio(io.ComfyNode):
|
||||
display_name="LongCat_Video_SM_Audio",
|
||||
category="LongCat_Video",
|
||||
inputs=[
|
||||
io.AudioEncoder.Input("audio_encoder"),
|
||||
io.Audio.Input("audio"),
|
||||
io.Int.Input("save_fps", default=25, min=8, max=1024, step=1),
|
||||
io.Int.Input("num_segments", default=1, min=1, max=1024, step=1),
|
||||
@@ -157,7 +158,7 @@ class LongCat_Video_SM_Audio(io.ComfyNode):
|
||||
],
|
||||
)
|
||||
@classmethod
|
||||
def execute(cls, audio,save_fps,num_segments,audio_type,p_box,left_audio=None) -> io.NodeOutput:
|
||||
def execute(cls, audio_encoder,audio,save_fps,num_segments,audio_type,p_box,left_audio=None) -> io.NodeOutput:
|
||||
if p_box:
|
||||
import ast
|
||||
# 将类似 "[100, 80, 800, 640], [1001, 80, 800, 640]" 的字符串转为嵌套列表
|
||||
@@ -165,7 +166,7 @@ class LongCat_Video_SM_Audio(io.ComfyNode):
|
||||
assert isinstance(parsed_p_box, list) and len(parsed_p_box) >= 2 , "p_box must be a list of int ,and must lens >2"
|
||||
else:
|
||||
parsed_p_box = None
|
||||
au_cond=get_audio_emb(weigths_longcat_current_path,audio,left_audio,audio_type,save_fps,num_segments,device,p_box=parsed_p_box)
|
||||
au_cond=get_audio_emb(audio_encoder,audio,left_audio,audio_type,save_fps,num_segments,device,p_box=parsed_p_box)
|
||||
clear_comfyui_cache()
|
||||
return io.NodeOutput(au_cond)
|
||||
|
||||
@@ -177,17 +178,37 @@ class LongCat_Video_SM_Vocal(io.ComfyNode):
|
||||
display_name="LongCat_Video_SM_Vocal",
|
||||
category="LongCat_Video",
|
||||
inputs=[
|
||||
io.AudioEncoder.Input("audio_encoder"),
|
||||
io.Audio.Input("audio"),
|
||||
io.Audio.Input("left_audio",optional=True),
|
||||
],
|
||||
outputs=[
|
||||
io.Audio.Output(display_name="audio"),
|
||||
io.Audio.Output(display_name="left_audio"),
|
||||
io.String.Output(display_name="audio_path"),
|
||||
],
|
||||
)
|
||||
@classmethod
|
||||
def execute(cls, audio, left_audio=None) -> io.NodeOutput:
|
||||
left_audio_path=audio2path(left_audio) if left_audio is not None else None
|
||||
audio_path,audio,left_audio=get_audio_vocal(weigths_longcat_current_path,audio2path(audio),left_audio_path,folder_paths.get_output_directory())
|
||||
return io.NodeOutput(audio,left_audio,audio_path)
|
||||
def execute(cls, audio_encoder,audio,) -> io.NodeOutput:
|
||||
audio_path,audio=get_audio_vocal(audio_encoder,audio2path(audio),folder_paths.get_output_directory())
|
||||
return io.NodeOutput(audio,audio_path)
|
||||
|
||||
class LongCat_Video_SM_VocalModel(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="LongCat_Video_SM_VocalModel",
|
||||
display_name="LongCat_Video_SM_VocalModel",
|
||||
category="LongCat_Video",
|
||||
inputs=[
|
||||
io.Combo.Input(
|
||||
"audio_encoder_vocal",options=["none"]+[i for i in folder_paths.get_filename_list("longcat") if i.endswith(".onnx")],
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.AudioEncoder.Output(),
|
||||
],
|
||||
)
|
||||
@classmethod
|
||||
def execute(cls, audio_encoder_vocal) -> io.NodeOutput:
|
||||
vocal_separator_path=folder_paths.get_full_path_or_raise("longcat", audio_encoder_vocal) if audio_encoder_vocal!="none" else None
|
||||
audio_encoder=load_audio_vocal(vocal_separator_path,folder_paths.get_output_directory(),weigths_longcat_current_path)
|
||||
return io.NodeOutput(audio_encoder)
|
||||
@@ -3,7 +3,8 @@
|
||||
|
||||
Update
|
||||
----
|
||||
* add node
|
||||
* use comfyUI whisper-large-v3 audio_encoders,offload lora to cpu
|
||||
* 音频解码改成comfyUI内置单体模型方式,开启lora和cache的卸载,降低显存占用
|
||||
* 部署节点,目前单人和双人测试通过,部分参数未严谨测试,请自行测试,有bug可以提issues或反馈给我的B站或小红书smthem账号
|
||||
|
||||
1.Installation
|
||||
@@ -33,11 +34,11 @@ links: [vae/vocal_separator/whisper-large-v3/lora](https://huggingface.co/meitua
|
||||
| ├── LongCat_Avatar_1.5_vae.safetensors
|
||||
├── ComfyUI/models/clip/
|
||||
| ├── umt5_xxl_fp8_e4m3fn_scaled.safetensors
|
||||
├── ComfyUI/models/longcat/ # 懒得写节点
|
||||
| ├── vocal_separator
|
||||
| ├── 4 files
|
||||
| ├── whisper-large-v3
|
||||
| ├── 13 files ,model.safetensors,config.json,tokenizer.json,vocab.json.. #模型只下载model.safetensors即可,json要下
|
||||
├── ComfyUI/models/audio_encoders/
|
||||
| ├── whisper-large-v3.safetensors # rename or not
|
||||
├── ComfyUI/models/longcat/
|
||||
| ├── Kim_Vocal_2.onnx # 配套config文件会自动下,可以下了先放进去
|
||||
|
||||
```
|
||||
|
||||
4 Example
|
||||
|
||||
+3
-1
@@ -1,7 +1,8 @@
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
from typing_extensions import override
|
||||
|
||||
from .LongCat_Video_node import LongCat_Video_SM_Model, LongCat_Video_SM_Sampler,LongCat_Video_SM_Encode,LongCat_Video_SM_Audio,LongCat_Video_SM_Vocal
|
||||
from .LongCat_Video_node import LongCat_Video_SM_Model, LongCat_Video_SM_Sampler,LongCat_Video_SM_Encode,LongCat_Video_SM_Audio,LongCat_Video_SM_Vocal,LongCat_Video_SM_VocalModel
|
||||
|
||||
class LongCat_Video_SM_Extension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
@@ -11,6 +12,7 @@ class LongCat_Video_SM_Extension(ComfyExtension):
|
||||
LongCat_Video_SM_Encode,
|
||||
LongCat_Video_SM_Audio,
|
||||
LongCat_Video_SM_Vocal,
|
||||
LongCat_Video_SM_VocalModel
|
||||
]
|
||||
|
||||
async def comfy_entrypoint() -> LongCat_Video_SM_Extension: # ComfyUI calls this to load your extension and its nodes.
|
||||
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 352 KiB After Width: | Height: | Size: 411 KiB |
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user