This commit is contained in:
smthemex
2026-02-10 21:16:35 +08:00
parent 026eb69707
commit 6dd8f162d7
11 changed files with 1641 additions and 841 deletions
+1 -1
View File
@@ -71,7 +71,7 @@ class InteractAvatar_SM_Predata(io.ComfyNode):
io.Vae.Input("vae"),
io.Image.Input("images"), # image or video
io.Image.Input("pose_images"), # image or video
io.Int.Input("short_side", default=512, min=256, max=nodes.MAX_RESOLUTION,step=32,display_mode=io.NumberDisplay.number),
io.Combo.Input("short_side",options= [704,512]),
io.String.Input("prompt",multiline=True, default="两只手打招呼 伸出大拇指点赞"),
io.String.Input("negative_prompt",multiline=True, default="bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"),
io.String.Input("structured_prompt",multiline=True, default=" (Raise both hands slowly) (Wave both hands side to side for greeting),\n (Make a fist) (Extend the thumb upwards)"),
+6 -1
View File
@@ -20,7 +20,10 @@
# ComfyUI_InteractAvatar
InteractAvatar is a novel dual-stream DiT framework that enables talking avatars to perform Grounded Human-Object Interaction (GHOI)
# Tips
# Update
* fix bug ,now output video short side muse be 512 or 704
* If your Vram <24G,turn on 'offload', ActionAndSong mode use 'long model' and need chocie '2' mode;example img\video\ audio in "InterDemo" dir
* test env 64G RAM, 12G VRAM,win11
* The prompt words for the singing mode and the action prompt words must have the same number of lines;
@@ -56,6 +59,8 @@ pip install -r requirements.txt
![](https://github.com/smthemex/ComfyUI_InteractAvatar/blob/main/example_workflows/example-song.png)
* object
![](https://github.com/smthemex/ComfyUI_InteractAvatar/blob/main/example_workflows/example.png)
* ap2v audio and pose driver
![](https://github.com/smthemex/ComfyUI_InteractAvatar/blob/main/example_workflows/example_ap2v.png)
# 5 Citation
```
File diff suppressed because it is too large Load Diff
Binary file not shown.

After

Width:  |  Height:  |  Size: 442 KiB

+32 -2
View File
@@ -13,7 +13,7 @@ import torchaudio
import folder_paths
from comfy.utils import common_upscale,ProgressBar
from safetensors.torch import load_file
import soundfile as sf
import comfy.model_management as mm
from pathlib import PureWindowsPath
cur_path = os.path.dirname(os.path.abspath(__file__))
@@ -136,7 +136,22 @@ def clear_comfyui_cache():
max_gpu_memory = torch.cuda.max_memory_allocated()
print(f"After Max GPU memory allocated: {max_gpu_memory / 1000 ** 3:.2f} GB")
# def trans2path(audio):
# if audio is None:
# return None
# import io as io_base
# audio_file_prefix = ''.join(random.choice("0123456789") for _ in range(6))
# audio_file = os.path.join(folder_paths.get_input_directory(), f"audio_{audio_file_prefix}_temp.wav")
# buff = io_base.BytesIO()
# torchaudio.save(buff, audio["waveform"].squeeze(0), audio["sample_rate"], format="FLAC")
# with open(audio_file, 'wb') as f:
# f.write(buff.getbuffer())
# return audio_file
def trans2path(audio):
"""
修正版:使用 soundfile 代替 torchaudio.save 以避开 torchcodec 的环境报错。
"""
if audio is None:
return None
import io as io_base
@@ -144,12 +159,27 @@ def trans2path(audio):
audio_file = os.path.join(folder_paths.get_input_directory(), f"audio_{audio_file_prefix}_temp.wav")
buff = io_base.BytesIO()
torchaudio.save(buff, audio["waveform"].squeeze(0), audio["sample_rate"], format="FLAC")
# --- 修正逻辑开始 ---
# ComfyUI 音频格式通常为 [Batch, Channels, Samples] -> [1, C, S]
# 我们需要将其转换为 NumPy,并调整维度为 soundfile 要求的 [Samples, Channels]
waveform = audio["waveform"].squeeze(0).cpu().numpy() # 结果为 [C, S]
sample_rate = audio["sample_rate"]
if waveform.ndim == 2:
waveform = waveform.T # 转置为 [S, C]
# 使用 soundfile 直接写入内存流,不触发 torchaudio 的后端检测
sf.write(buff, waveform, sample_rate, format="FLAC")
# --- 修正逻辑结束 ---
with open(audio_file, 'wb') as f:
f.write(buff.getbuffer())
return audio_file
def encode_image( image, vae):
if image is None:
return None
+14 -11
View File
@@ -269,7 +269,6 @@ def perdata( clip,vae,images,dw_iamges,object_images,object_mask,audio_path,mode
dw_seqs = [dw.resize(dw_img.size, Image.LANCZOS) for dw in dw_seqs]
dwpose_len = len(dw_seqs)
dwpose_frame_num = (dwpose_len - 1) // 4 * 4 + 1
# pre audio
if audio_path is not None:
audio_input, sampling_rate = librosa.load(audio_path, sr=16000)
@@ -281,7 +280,7 @@ def perdata( clip,vae,images,dw_iamges,object_images,object_mask,audio_path,mode
audio_frame_num = dwpose_frame_num
sampling_rate = 16000
if dw_seqs is not None:
if dw_seqs is not None: # 对齐pose和音频帧数
dw_seqs = dw_seqs[:frame_num]
if audio_frame_num > dwpose_frame_num:
padding_dwpose = dw_seqs[-1]
@@ -292,9 +291,10 @@ def perdata( clip,vae,images,dw_iamges,object_images,object_mask,audio_path,mode
audio_frames_clip = np.concatenate([audio_frames_clip, padding_audio], axis=0)
audio_frame_num = dwpose_frame_num
if dw_seqs is None and mode == 'ap2v':
dw_seqs = [dw_img] * audio_frame_num
dw_seqs = [dw_img] * audio_frame_num #对齐帧数
frame_num = min(frame_num,min(audio_frame_num,dwpose_frame_num)) #对齐推理帧和音频帧数
frame_num = min(frame_num,min(audio_frame_num,dwpose_frame_num))
audio_frames_clip = audio_frames_clip[:int(frame_num * sampling_rate / 25)]
wav2vec_feature_extractor, audio_encoder= custom_init(device, wav2vec_dir)
@@ -330,8 +330,8 @@ def perdata( clip,vae,images,dw_iamges,object_images,object_mask,audio_path,mode
if mode in ['a2v','a2mv','mv','i2v']:
dw_seqs = None
vae_stride=WAN_CONFIGS['ti2v-5B'].vae_stride
patch_size=WAN_CONFIGS['ti2v-5B'].patch_size
vae_stride=WAN_CONFIGS['ti2v-5B'].vae_stride #(4, 16, 16)
patch_size=WAN_CONFIGS['ti2v-5B'].patch_size #(1, 2, 2)
sp_size=1
if back_append_frame==1:
@@ -345,15 +345,16 @@ def perdata( clip,vae,images,dw_iamges,object_images,object_mask,audio_path,mode
if dw_seqs is not None:
if isinstance(dw_seqs, list): # If input is a list of PIL Images
processed_poses = [phi2narry(p) for p in dw_seqs]
cond_pose_sequence = torch.stack(processed_poses).to(device)
processed_poses = [phi2narry(p.convert('RGB')) for p in dw_seqs]
cond_pose_sequence=torch.cat(processed_poses).to(device)
#cond_pose_sequence = torch.stack(processed_poses).to(device)
else: # If input is already a tensor
cond_pose_sequence = dw_seqs.to(device)
dwpose_len = cond_pose_sequence.shape[0]
dwpose_len = (dwpose_len - 1) // vae_stride[0] * vae_stride[0] + 1
frame_num = min(frame_num, dwpose_len)
cond_pose_sequence = cond_pose_sequence[:frame_num]
#print(cond_pose_sequence.shape) #torch.Size([133, 256, 448, 3])
else:
cond_pose_sequence = torch.ones((frame_num, pose_ref_img.shape[1], pose_ref_img.shape[2], pose_ref_img.shape[3]), device=pose_ref_img.device, dtype=pose_ref_img.dtype) * 0.5
@@ -419,6 +420,7 @@ def perdata( clip,vae,images,dw_iamges,object_images,object_mask,audio_path,mode
mode=mode,
frame_num=frame_num,
max_frames_num=frame_num,
short_side=short_side,
)
else:
curr_cond_image=phi2narry(img) #BHWC
@@ -434,8 +436,8 @@ def perdata( clip,vae,images,dw_iamges,object_images,object_mask,audio_path,mode
# Pose Sequence 处理
if dw_seqs is not None:
if isinstance(dw_seqs, list):
processed_poses = [phi2narry(p) for p in dw_seqs]
cond_pose_sequence_full = torch.stack(processed_poses).to(device)
processed_poses = [phi2narry(p.convert('RGB')) for p in dw_seqs]
cond_pose_sequence=torch.cat(processed_poses).to(device)
else:
cond_pose_sequence_full = dw_seqs.to(device)
else:
@@ -535,5 +537,6 @@ def perdata( clip,vae,images,dw_iamges,object_images,object_mask,audio_path,mode
frame_num=frame_num,
vae=vae,
max_frames_num=frame_num,
short_side=short_side,
)
return data_dict
+10 -26
View File
@@ -1479,6 +1479,7 @@ class WanModel(ModelMixin, ConfigMixin):
use_gradient_checkpointing_offload=False,
cond_flag=False,
gpu_manager=None,
up_scale=2.75,
**kwargs
):
r"""
@@ -1512,16 +1513,14 @@ class WanModel(ModelMixin, ConfigMixin):
device = self.patch_embedding.weight.device
if self.freqs.device != device:
self.freqs = self.freqs.to(device)
# print(x.shape)
# embeddings
if x is not None:
x = [self.patch_embedding(u.unsqueeze(0)) for u in x] # [b, 1, dim, t, h/2, w/2] -> [b,seq_len,dim 1536]
grid_sizes = torch.stack(
[torch.tensor(u.shape[2:], dtype=torch.long) for u in x]) # [1, 3], thw
frame_l = x[0].shape[-1] * x[0].shape[-2]
x = [u.flatten(2).transpose(1, 2) for u in x] # [b, dim, thw/4] => [b, thw/4, dim]
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
#print(seq_lens,"seq lens") #tensor([10368]) seq lens
if seq_len!=0:
assert seq_lens.max() <= seq_len
x = torch.cat([
@@ -1672,7 +1671,6 @@ class WanModel(ModelMixin, ConfigMixin):
)
iii = 0
pre_motion = torch.zeros_like(motion, requires_grad=True)
#print(f"len(self.blocks): {len(self.blocks)}",len(self.zero_motion_proj_blocks),len(self.motion_blocks)) #30,30,30
for idx,block in enumerate(self.blocks):
if gpu_manager is not None:
if idx < len(self.blocks):
@@ -1690,11 +1688,12 @@ class WanModel(ModelMixin, ConfigMixin):
residual_motion = residual_motion.transpose(1, 2).reshape(bb, -1, ff, hh, ww)
residual_motion = residual_motion.permute(0,2,1,3,4).reshape(bb*ff, -1, hh, ww)
residual_motion_up = F.interpolate(
residual_motion, scale_factor=(2.75, 2.75),
residual_motion, scale_factor=(up_scale, up_scale), #2.75 = 704/256,2.0=512/256
mode='bilinear', align_corners=False
).to(motion.dtype)
H_out, W_out = residual_motion_up.shape[-2], residual_motion_up.shape[-1]
last_part_len = H_out * W_out
residual_motion_up = residual_motion_up.reshape(bb, ff, -1, H_out, W_out).permute(0,2,1,3,4)
if gpu_manager is not None:
@@ -1706,29 +1705,14 @@ class WanModel(ModelMixin, ConfigMixin):
prev_module_ = gpu_manager.managed_motion_proj_modules[idx - 1]
if hasattr(prev_module_, 'to'):
prev_module_.to('cpu')
value_to_add = self.zero_motion_proj_blocks[idx](residual_motion_up.flatten(2).transpose(1, 2))
#print(f"value_to_add shape: {value_to_add.shape}",f"x shape: {x.shape}",f"last_part_len: {last_part_len}") #value_to_add shape: torch.Size([1, 15972, 3072])
#print(f"value_to_add shape: {value_to_add.shape}",f"x shape: {x.shape}",f"last_part_len: {last_part_len}")
#value_to_add shape: torch.Size([1, 19602, 3072]) x shape: torch.Size([1, 10368, 3072]) last_part_len: 726 # if use short_side=512 got error when upscale=2.75
x[:, :-last_part_len, :] = x[:, :-last_part_len, :] + value_to_add[:, :-last_part_len, :]
#x[:, :-last_part_len, :] = x[:, :-last_part_len, :] + value_to_add[:, :-last_part_len, :]
# 确保 last_part_len 不超过 x 的序列长度
last_part_len = min(last_part_len, x.size(1))
# 计算 x 中需要更新的长度
x_update_len = x.size(1) - last_part_len
# 确保 value_to_add 的长度与 x_update_len 匹配
if value_to_add.size(1) >= x_update_len:
# 如果 value_to_add 足够长,只使用前 x_update_len 个位置
x[:, :x_update_len, :] = x[:, :x_update_len, :] + value_to_add[:, :x_update_len, :]
else:
# 如果 value_to_add 比 x_update_len 短,使用全部 value_to_add
x[:, :value_to_add.size(1), :] = x[:, :value_to_add.size(1), :] + value_to_add
## apply motion block
if gpu_manager is not None:
if idx < len(self.motion_blocks):
module_1 = gpu_manager.managed_motion_modules[idx]
+25 -19
View File
@@ -161,7 +161,8 @@ class WanTIA2MVRefBackIDPrefix:
)
logging.info(f"Creating WanModel from {checkpoint_dir}")
model_cfg = config
model_type = 'ti2v' if model_cfg.__name__ == 'Config: Wan TI2V 5B' else 'i2v'
#print(model_cfg.__name__)
model_type = 'ti2v' if model_cfg.__name__ == 'Config: Wan TI2V 5B' else 'i2v' #Config: Wan TI2V 5B
ctx = init_empty_weights if is_accelerate_available() else nullcontext
with ctx():
unet = WanModel(
@@ -298,7 +299,8 @@ class WanTIA2MVRefBackIDPrefix:
- H: Frame height (from max_area)
- W: Frame width from max_area)
"""
up_scale = 2.75 if kwargs.get('short_side', 704)==704 else 2.0
#print(f'frame_num: {frame_num}')
if self.origin_mode:
cond_image = TF.to_tensor(img).sub_(0.5).div_(0.5).to(self.device)
cond_image = cond_image[None, :, None, :, :]
@@ -398,9 +400,10 @@ class WanTIA2MVRefBackIDPrefix:
if self.origin_mode:
if mode in ['a2v','a2mv','mv','i2v']:
cond_pose_sequence = torch.zeros_like(cond_pose_sequence)
#print(f'mode: {mode}')
if mode in ['p2v','mv','i2v']:
audio_embs = zero_audio_embs
if self.origin_mode:
h, w = cond_image.shape[-2], cond_image.shape[-1]
lat_h, lat_w = h // self.vae_stride[1], w // self.vae_stride[2]
@@ -411,12 +414,13 @@ class WanTIA2MVRefBackIDPrefix:
lat_h=kwargs.get('lat_h', None)
lat_w=kwargs.get('lat_w', None)
max_seq_len=kwargs.get('max_seq_len', None)
#print(f'max_seq_len: {max_seq_len}',lat_h,lat_w) # 15232 32 56 34*32/2*56/2= 15232
noise = torch.randn(
1, 48, (frame_num - 1) // 4 + 1 + 1,
lat_h,
lat_w,
dtype=torch.float32,
device=self.device)
device=self.device) # 初始噪声加多4帧
if self.origin_mode:
cond_image = self.vae.encode([cond_image.squeeze(0).to(torch.float32)])[0].unsqueeze(0)
obj_image = self.vae.encode([obj_image.repeat(1, 1, 4, 1, 1).squeeze(0).to(torch.float32)])[0].unsqueeze(0)
@@ -437,7 +441,7 @@ class WanTIA2MVRefBackIDPrefix:
motion_lat_h=kwargs.get('motion_lat_h', None)
motion_lat_w=kwargs.get('motion_lat_w', None)
motion_max_seq_len=kwargs.get('motion_max_seq_len', None)
#print(f'motion_max_seq_len: {motion_max_seq_len}',motion_lat_h,motion_lat_w) #motion_max_seq_len: 2496 24 16
motion_noise = torch.randn(
1, 48, (frame_num - 1) // 4 + 1 + 1,
motion_lat_h,
@@ -479,14 +483,15 @@ class WanTIA2MVRefBackIDPrefix:
else:
gpu_manager = BlockGPUManager(device="cuda")
gpu_manager.setup_for_inference(self.model)
#print(latent.shape,cond_image.shape,obj_image.shape) #torch.Size([1, 48, 22, 48, 32]) torch.Size([1, 48, 1, 48, 32]) torch.Size([1, 48, 1, 48, 32])
#print(latent.shape,cond_image.shape,obj_image.shape) #torch.Size([1, 48, 35, 32, 56]) torch.Size([1, 48, 1, 32, 56]) torch.Size([1, 48, 1, 32, 56])
latent[:, :, 0:1] = cond_image
latent[:, :, -1:] = obj_image
#print(latent_motion.shape,cond_dw_img.shape,small_img.shape) #torch.Size([1, 48, 22, 24, 16]) torch.Size([1, 48, 1, 24, 16]) torch.Size([1, 48, 1, 24, 16])
#print(latent_motion.shape,cond_dw_img.shape,small_img.shape) #torch.Size([1, 48, 35, 16, 28]) torch.Size([1, 48, 1, 16, 28]) torch.Size([1, 48, 1, 16, 28])
latent_motion[:, :, 0:1] = cond_dw_img
latent_motion[:, :, -1:] = small_img
#print(cond_pose_sequence.shape) #torch.Size([1, 48, 21, 24, 16])
cond_pose_sequence = torch.concat([cond_pose_sequence, small_img], dim=2)
#print(f'cond_pose_sequence.shape: {cond_pose_sequence.shape}') #torch.Size([1, 48, 35, 16, 28])
cond_pose_sequence = cond_pose_sequence.to(latent_motion.dtype).to(self.device,cur_dtype)
audio_embs = audio_embs.to(latent_motion.dtype).to(self.device,cur_dtype)
# zero_audio_embs = zero_audio_embs.to(latent_motion.dtype).to(self.device)
@@ -494,8 +499,8 @@ class WanTIA2MVRefBackIDPrefix:
# prepare condition and uncondition configs
arg_c = {
'context_list': [context[0].to(cur_dtype)],
'seq_len': max_seq_len + lat_h * lat_w //4,
'seq_len_motion': motion_max_seq_len + motion_lat_h * motion_lat_w //4,
'seq_len': max_seq_len + lat_h * lat_w //4, # 补齐4帧长度
'seq_len_motion': motion_max_seq_len + motion_lat_h * motion_lat_w //4, # 15232+
'mode': mode,
'skip_block': False,
'audio_embedding': audio_embs,
@@ -531,11 +536,11 @@ class WanTIA2MVRefBackIDPrefix:
arg_null_text['skip_block'] = False
#print('no bad_cfg', timestep, bad_cfg)
if mode not in ['mv','a2mv']:
noise_pred_cond, noise_pred_cond_motion = self.model(
x=latent_model_input,motion=cond_pose_sequence,
t=timestep,motion_t=zero_timestep, **arg_c,gpu_manager=gpu_manager,
t=timestep,motion_t=zero_timestep, **arg_c,gpu_manager=gpu_manager,up_scale=up_scale,
)
torch_gc()
noise_pred_drop_text, noise_pred_drop_text_motion = self.model(
@@ -543,14 +548,14 @@ class WanTIA2MVRefBackIDPrefix:
motion=cond_pose_sequence,
t=timestep,motion_t=zero_timestep,
**arg_null_text,
gpu_manager=gpu_manager,
gpu_manager=gpu_manager,up_scale=up_scale,
)
torch_gc()
else:
# inference with CFG strategy
noise_pred_cond, noise_pred_cond_motion = self.model(
x=latent_model_input,motion=latent_model_input_motion,
t=timestep,motion_t=timestep, **arg_c,gpu_manager=gpu_manager,
t=timestep,motion_t=timestep, **arg_c,gpu_manager=gpu_manager,up_scale=up_scale,
)
torch_gc()
noise_pred_drop_text, noise_pred_drop_text_motion = self.model(
@@ -559,7 +564,7 @@ class WanTIA2MVRefBackIDPrefix:
t=timestep,
motion_t=timestep,
**arg_null_text,
gpu_manager=gpu_manager,
gpu_manager=gpu_manager,up_scale=up_scale,
)
torch_gc()
@@ -649,6 +654,7 @@ class WanTIA2MVRefBackIDPrefix:
Generates video frames autoregressively with Latent Prefix Loop.
Segment Length: 101 frames (26 temporal latents + 2 condition latents).
"""
up_scale = 2.75 if kwargs.get('short_side', 704)==704 else 2.0
if self.origin_mode:
# --- 1. 输入预处理 (Input Preprocessing) ---
# 初始首帧 (Start Frame) - 仅用于第一段
@@ -996,14 +1002,14 @@ class WanTIA2MVRefBackIDPrefix:
# Model Forward
if three_cfg:
noise_pred_cond, noise_pred_cond_motion = self.model(
x=latent_model_input, motion=motion_input, t=timestep, motion_t=motion_time, **arg_c,gpu_manager=gpu_manager)
x=latent_model_input, motion=motion_input, t=timestep, motion_t=motion_time, **arg_c,gpu_manager=gpu_manager,up_scale=up_scale,)
torch_gc()
noise_pred_drop_text, noise_pred_drop_text_motion = self.model(
x=latent_model_input,
motion=motion_input,
t=timestep,
motion_t=motion_time,
**arg_null_text,gpu_manager=gpu_manager
**arg_null_text,gpu_manager=gpu_manager,up_scale=up_scale
)
torch_gc()
noise_pred_pure_audio, noise_pred_pure_audio_motion = self.model(
@@ -1011,7 +1017,7 @@ class WanTIA2MVRefBackIDPrefix:
motion=motion_input,
t=timestep,
motion_t=motion_time,
**arg_pure_audio,gpu_manager=gpu_manager
**arg_pure_audio,gpu_manager=gpu_manager,up_scale=up_scale,
)
torch_gc()
@@ -1023,11 +1029,11 @@ class WanTIA2MVRefBackIDPrefix:
)
else:
noise_pred_cond, noise_pred_cond_motion = self.model(
x=latent_model_input, motion=motion_input, t=timestep, motion_t=motion_time, **arg_c,gpu_manager=gpu_manager)
x=latent_model_input, motion=motion_input, t=timestep, motion_t=motion_time, **arg_c,gpu_manager=gpu_manager,up_scale=up_scale,)
torch_gc()
noise_pred_drop_text, noise_pred_drop_text_motion = self.model(
x=latent_model_input, motion=motion_input, t=timestep,
motion_t=motion_time, **arg_null_text,gpu_manager=gpu_manager)
motion_t=motion_time, **arg_null_text,gpu_manager=gpu_manager,up_scale=up_scale,)
torch_gc()
noise_pred = noise_pred_drop_text + text_guide_scale * (noise_pred_cond - noise_pred_drop_text)