init
This commit is contained in:
@@ -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)"),
|
||||
|
||||
@@ -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
|
||||

|
||||
* object
|
||||

|
||||
* ap2v audio and pose driver
|
||||

|
||||
|
||||
# 5 Citation
|
||||
```
|
||||
|
||||
+1553
-781
File diff suppressed because it is too large
Load Diff
Binary file not shown.
|
After Width: | Height: | Size: 442 KiB |
+32
-2
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user