Merge pull request #3 from smthemex/Pr

init
This commit is contained in:
smthemex
2026-05-28 13:37:51 +08:00
committed by GitHub
12 changed files with 794 additions and 812 deletions
+163 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+7 -6
View File
@@ -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
View File
@@ -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