feat: Add VAD-based silence handling and Blink rate control for Ditto

This commit is contained in:
Hawk Lee
2026-01-21 01:46:04 +08:00
parent 7dddbbb0e0
commit cd0a8d23b5
3 changed files with 64 additions and 6 deletions
+50 -2
View File
@@ -416,6 +416,9 @@ class AIIA_DittoSampler:
"hd_rot_y": ("FLOAT", {"default": 0.0, "min": -30.0, "max": 30.0, "step": 1.0}),
"hd_rot_r": ("FLOAT", {"default": 0.0, "min": -30.0, "max": 30.0, "step": 1.0}),
"mouth_amp": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.05}),
"auto_ref_frame": ("BOOLEAN", {"default": True, "label_on": "Auto-Revert on Silence", "label_off": "Disabled"}),
"ref_threshold": ("FLOAT", {"default": 0.01, "min": 0.0, "max": 1.0, "step": 0.001}),
"blink_mode": (["Random (Normal)", "Fast", "Slow", "None"], {"default": "Random (Normal)"}),
}
}
@@ -424,7 +427,7 @@ class AIIA_DittoSampler:
FUNCTION = "generate"
CATEGORY = "AIIA/Ditto"
def generate(self, pipe, ref_image, audio, sampling_steps, fps, crop_scale, emo, drive_eye, chk_eye_blink, smo_k_d, hd_rot_p, hd_rot_y, hd_rot_r, mouth_amp):
def generate(self, pipe, ref_image, audio, sampling_steps, fps, crop_scale, emo, drive_eye, chk_eye_blink, smo_k_d, hd_rot_p, hd_rot_y, hd_rot_r, mouth_amp, auto_ref_frame, ref_threshold, blink_mode):
# pipe is the dict we returned in Loader
master_sdk = pipe["sdk"]
cfg_pkl = pipe["cfg_pkl"]
@@ -453,6 +456,48 @@ class AIIA_DittoSampler:
target_fps = 25 # Force 25 for stability first
num_frames = math.ceil(len(audio_np) / 16000 * target_fps)
# VAD / Volume Analysis
ctrl_info = {}
if auto_ref_frame:
# Calculate RMS per frame
frame_len = 640 # 16000 / 25
# Simple RMS calculation
# Pad audio if needed
pad_len = num_frames * frame_len - len(audio_np)
if pad_len > 0:
audio_proc = np.pad(audio_np, (0, pad_len))
else:
audio_proc = audio_np[:num_frames * frame_len]
# Reshape to (num_frames, frame_len)
audio_frames = audio_proc.reshape(num_frames, frame_len)
rms = np.sqrt(np.mean(audio_frames**2, axis=1))
for i in range(num_frames):
val = rms[i]
if val < ref_threshold:
alpha = 0.0 # Full Reference
else:
alpha = 1.0 # Full Generated
if alpha < 1.0:
ctrl_info[i] = {"vad_alpha": alpha}
# Blink Settings
delta_eye_open_n = 0 if chk_eye_blink else -1
blink_min = 60
blink_max = 100
if blink_mode == "Fast":
blink_min = 10
blink_max = 40
elif blink_mode == "Slow":
blink_min = 120
blink_max = 200
elif blink_mode == "None":
delta_eye_open_n = -1
# Map emo string to int
emo_map = {
"Angry": 0, "Disgust": 1, "Fear": 2, "Happy": 3,
@@ -479,9 +524,12 @@ class AIIA_DittoSampler:
crop_scale=crop_scale,
emo=emo_idx,
drive_eye=drive_eye,
delta_eye_open_n=0 if chk_eye_blink else -1, # 0=random, -1=none
delta_eye_open_n=delta_eye_open_n, # 0=random, -1=none
blink_interval_min=blink_min,
blink_interval_max=blink_max,
smo_k_d=smo_k_d,
overall_ctrl_info=overall_ctrl_info,
ctrl_info=ctrl_info,
total_frames=num_frames # For pbar
)
@@ -100,7 +100,7 @@ def _mix_s_d_info(
return x_d_info
def _set_eye_blink_idx(N, blink_n=15, open_n=-1):
def _set_eye_blink_idx(N, blink_n=15, open_n=-1, interval_min=60, interval_max=100):
"""
open_n:
-1: no blink
@@ -108,8 +108,8 @@ def _set_eye_blink_idx(N, blink_n=15, open_n=-1):
>0: fix open_n
list: loop open_n
"""
OPEN_MIN = 60
OPEN_MAX = 100
OPEN_MIN = interval_min
OPEN_MAX = interval_max
idx = [0] * N
if isinstance(open_n, int):
@@ -311,6 +311,8 @@ class MotionStitch:
drive_eye=None, # use d eye or s eye
delta_eye_arr=None, # fix eye
delta_eye_open_n=-1, # int|list
blink_interval_min=60,
blink_interval_max=100,
fade_out_keys=("exp",),
fade_type="", # "" | "d0" | "s"
flag_stitching=True,
@@ -366,7 +368,11 @@ class MotionStitch:
if self.drive_eye and self.delta_eye_arr is not None:
N = 3000 if self.N_d == -1 else self.N_d
self.delta_eye_idx_list = _set_eye_blink_idx(
N, len(self.delta_eye_arr), self.delta_eye_open_n
N,
len(self.delta_eye_arr),
self.delta_eye_open_n,
interval_min=blink_interval_min,
interval_max=blink_interval_max
)
self.pose_s = None
+4
View File
@@ -117,6 +117,8 @@ class StreamSDK:
self.drive_eye = kwargs.get("drive_eye", None) # None: true4image, false4video
self.delta_eye_arr = kwargs.get("delta_eye_arr", None)
self.delta_eye_open_n = kwargs.get("delta_eye_open_n", 0)
self.blink_interval_min = kwargs.get("blink_interval_min", 60)
self.blink_interval_max = kwargs.get("blink_interval_max", 100)
self.fade_type = kwargs.get("fade_type", "") # "" | "d0" | "s"
self.fade_out_keys = kwargs.get("fade_out_keys", ("exp",))
self.flag_stitching = kwargs.get("flag_stitching", True)
@@ -188,6 +190,8 @@ class StreamSDK:
drive_eye=self.drive_eye,
delta_eye_arr=self.delta_eye_arr,
delta_eye_open_n=self.delta_eye_open_n,
blink_interval_min=self.blink_interval_min,
blink_interval_max=self.blink_interval_max,
fade_out_keys=self.fade_out_keys,
fade_type=self.fade_type,
flag_stitching=self.flag_stitching,