Align with the official FlashHead implementation (#488)

- Update predict_s2v.py with improved inference configuration
- Align pipeline_flashhead.py with official streaming, scheduling, conditioning, and motion-cache behavior
This commit is contained in:
Tyler Luan
2026-05-08 14:04:09 +08:00
committed by GitHub
parent 0266bab98b
commit 8fb0bb165c
2 changed files with 53 additions and 30 deletions
+2 -1
View File
@@ -96,7 +96,7 @@ color_correction_strength = 1.0
use_apg = False
apg_momentum = 0.5
apg_norm_threshold = 1.0
audio_encode_mode = "once"
audio_encode_mode = "stream"
device = set_multi_gpus_devices(ulysses_degree, ring_degree)
config = OmegaConf.load(config_path)
@@ -208,6 +208,7 @@ with torch.no_grad():
# Prepare ref_image latent for FlashHead (no clip_image needed)
ref_image = get_image_latent(ref_image, sample_size=sample_size)
ref_image = ref_image * 2.0 - 1.0
sample = pipeline(
segment_frame_length = segment_frame_length,
+51 -29
View File
@@ -147,6 +147,16 @@ def match_and_blend_colors(videos, original_color_reference, strength=0.5):
return videos
def official_flashhead_timesteps(num_inference_steps, shift, device, num_timesteps=1000):
if num_inference_steps == 4:
timesteps = [1000, 750, 500, 250]
else:
timesteps = np.linspace(num_timesteps, 1, num_inference_steps, dtype=np.float32).tolist()
timesteps = torch.tensor(timesteps + [0.0], dtype=torch.float32, device=device)
t = timesteps / num_timesteps
return shift * t / (1 + (shift - 1) * t) * num_timesteps
@dataclass
class WanPipelineOutput(BaseOutput):
r"""
@@ -320,7 +330,7 @@ class FlashHeadPipeline(DiffusionPipeline):
segment_frame_length: int = 33,
num_inference_steps: int = 4,
timesteps: Optional[List[int]] = None,
audio_guide_scale: float = 4.0,
audio_guide_scale: float = 1.0,
num_videos_per_prompt: int = 1,
eta: float = 0.0,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
@@ -418,7 +428,7 @@ class FlashHeadPipeline(DiffusionPipeline):
tgt_fps = 25
slice_len = segment_frame_length - self.motion_frames # e.g. 33 - 5 = 28
human_speech_array_slice_len = slice_len * sample_rate // tgt_fps
cached_audio_duration = 2 # seconds (official default)
cached_audio_duration = 8 # seconds (official default)
cached_audio_length_sum = sample_rate * cached_audio_duration
# Pad audio to avoid truncating
@@ -475,7 +485,7 @@ class FlashHeadPipeline(DiffusionPipeline):
while True:
# 6. Prepare timesteps
if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, None, mu=1)
timesteps = official_flashhead_timesteps(num_inference_steps, shift, device)
elif isinstance(self.scheduler, FlowUniPCMultistepScheduler):
self.scheduler.set_timesteps(num_inference_steps, device=device, shift=shift)
timesteps = self.scheduler.timesteps
@@ -576,9 +586,8 @@ class FlashHeadPipeline(DiffusionPipeline):
lat_h_local = height // self.vae.spatial_compression_ratio
lat_w_local = width // self.vae.spatial_compression_ratio
# Encode condition video (ref_image in pixel space + zero frames)
cond_video_pixels = torch.zeros(1, 3, segment_frame_length, height, width, dtype=weight_dtype, device=device)
cond_video_pixels[:, :, :1] = ref_image[:, :, :1]
# Official FlashHead encodes the reference image repeated across the full clip.
cond_video_pixels = ref_image[:, :, :1].repeat(1, 1, segment_frame_length, 1, 1).to(device=device, dtype=weight_dtype)
y_latent_output = self.vae.encode(cond_video_pixels)
y = y_latent_output[0].mode().to(weight_dtype) # [B, 16, T_latent, H_latent, W_latent]
@@ -598,7 +607,8 @@ class FlashHeadPipeline(DiffusionPipeline):
self.transformer.num_inference_steps = num_inference_steps
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(tqdm(timesteps) if progress else timesteps):
denoise_timesteps = timesteps[:-1] if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler) else timesteps
for i, t in enumerate(tqdm(denoise_timesteps) if progress else denoise_timesteps):
self.transformer.current_steps = i
if self.interrupt:
@@ -628,28 +638,39 @@ class FlashHeadPipeline(DiffusionPipeline):
audio_wav2vec_fea=audio_emb_input,
)
# Audio guidance only (FlashHead does not use text CFG)
with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=device):
noise_pred_drop_audio = self.transformer(
x=latents,
t=timestep,
seq_len=seq_len,
y=y,
audio_wav2vec_fea=audio_emb_input * 0,
)
if use_apg:
diff_uncond_audio = noise_pred_cond - noise_pred_drop_audio
noise_pred = noise_pred_cond + (audio_guide_scale - 1) * adaptive_projected_guidance(
diff_uncond_audio, noise_pred_cond,
momentum_buffer=audio_momentumbuffer,
norm_threshold=apg_norm_threshold
)
# Official FlashHead Pro runs without classifier-free guidance.
if audio_guide_scale == 1.0 and not use_apg:
noise_pred = noise_pred_cond
else:
noise_pred = noise_pred_drop_audio + audio_guide_scale * (noise_pred_cond - noise_pred_drop_audio)
with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=device):
noise_pred_drop_audio = self.transformer(
x=latents,
t=timestep,
seq_len=seq_len,
y=y,
audio_wav2vec_fea=audio_emb_input * 0,
)
if use_apg:
diff_uncond_audio = noise_pred_cond - noise_pred_drop_audio
noise_pred = noise_pred_cond + (audio_guide_scale - 1) * adaptive_projected_guidance(
diff_uncond_audio, noise_pred_cond,
momentum_buffer=audio_momentumbuffer,
norm_threshold=apg_norm_threshold
)
else:
noise_pred = noise_pred_drop_audio + audio_guide_scale * (noise_pred_cond - noise_pred_drop_audio)
# compute the previous noisy sample x_t -> x_t-1
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
t_i = (timesteps[i] / self.scheduler.config.num_train_timesteps).to(weight_dtype)
t_i_1 = (timesteps[i + 1] / self.scheduler.config.num_train_timesteps).to(weight_dtype)
x_0 = latents - noise_pred * t_i
latents = (1 - t_i_1) * x_0 + t_i_1 * torch.randn(
latents.shape, dtype=latents.dtype, device=device, generator=generator
)
else:
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
# Always inject clean motion at the end (following official Line 773)
latents[:, :, :lat_motion_frames] = motion_latents[:, :, :lat_motion_frames]
@@ -661,7 +682,7 @@ class FlashHeadPipeline(DiffusionPipeline):
latents = callback_outputs.pop("latents", latents)
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
if i == len(denoise_timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
progress_bar.update()
if comfyui_progressbar:
pbar.update(1)
@@ -674,8 +695,9 @@ class FlashHeadPipeline(DiffusionPipeline):
if color_correction_strength > 0.0 and original_color_reference is not None:
image = match_and_blend_colors(image, original_color_reference, color_correction_strength)
# Cache generated video (following official Line 791-794)
if is_first_clip:
# Cache generated video. Official stream mode drops motion frames for every chunk,
# including the first chunk, because they are warm-up/cache frames.
if is_first_clip and audio_encode_mode != "stream":
gen_video_list.append(image)
else:
gen_video_list.append(image[:, :, cur_motion_frames_num:])