代码的角度分析Causal Forcing
一、项目概述与核心思想
论文思想解读见:彻底搞懂 Causal Forcing:实时交互式视频生成的因果蒸馏新范式-CSDN博客
1.1 研究背景与动机
Causal Forcing 是清华大学团队提出的自回归视频生成框架,其核心论文《Causal Forcing: Autoregressive Diffusion Distillation Done Right for High-Quality Real-Time Interactive Video Generation》发表于2026年2月。该工作要解决的核心问题是:如何在保持实时推理速度的同时,实现高质量的自回归视频生成?
传统的双向扩散模型虽然生成质量高,但无法实现流式生成——每一帧的生成都需要看到整个视频序列。自回归生成虽然天然支持流式输出,但面临严重的误差累积问题:早期帧的生成误差会随着时间传播放大,导致后续帧质量急剧下降。已有的方法如CausVid和Self Forcing虽然尝试解决这一问题,但在视觉质量和动态表现力上仍有不足。
Causal Forcing的核心洞察是:自回归视频生成中的误差累积,本质上源于"条件分布的不匹配"。在训练时,模型看到的是干净的真实帧作为条件;而在推理时,条件变成了模型自己生成的带有误差的帧。这种分布偏移(exposure bias)在多步扩散采样中被进一步放大。
1.2 核心技术贡献
论文提出了一个三阶段渐进式蒸馏框架,其核心思想可以概括为:
- 阶段一(AR Diffusion Training):训练一个因果结构的扩散模型,使模型学会"在给定历史帧条件下生成当前帧"。关键创新在于使用Teacher Forcing策略——训练时用真实的干净帧作为条件,而非自己生成的帧,避免了训练时的误差累积。
- 阶段二:将多步扩散模型压缩为少步模型。论文提供了两条路径:ODE Distillation(通过预计算ODE轨迹进行监督学习)和Consistency Distillation(无需预计算数据,直接在GT数据上训练)。这一阶段的核心目的是让模型学会用少量步数准确预测去噪方向。
- 阶段三:使用Distribution Matching Distillation完成最终蒸馏。此时教师模型是一个更强的双向扩散模型(14B参数),学生模型是因果模型(1.3B参数)。DMD通过对抗式训练,让学生模型的输出分布逼近教师模型的分布。
二、数据流程分析
2.1 数据的来源与组织形式
项目数据采用LMDB数据库存储,这是一种高效的键值存储格式,特别适合大规模视频latent的快速读取。数据流程的设计体现了论文对训练效率的考量——预计算好的latent避免了训练时重复进行VAE编解码。
数据存储位置:
- 原始GT数据:/Causal-Forcing/dataset/clean_data(LMDB格式)
- ODE轨迹数据:/Causal-Forcing/dataset/ODE6KCausal_framewise
数据加载核心实现:/Causal-Forcing/utils/dataset.py
# 数据集类的设计体现了不同训练阶段的需求
class ODERegressionLMDBDataset(Dataset):
"""用于Stage 2 ODE回归训练,输出完整的ODE轨迹"""
def __getitem__(self, idx):
return {
"prompts": prompts, # 文本描述
"ode_latent": torch.tensor(latents) # [num_steps, F, C, H, W] 从噪声到干净的完整轨迹
}
class LatentLMDBDataset(Dataset):
"""用于Stage 1和CD训练,输出干净的latent"""
def __getitem__(self, idx):
return {
"prompts": prompts,
"clean_latent": torch.tensor(latents) # [F, C, H, W] 干净的VAE latent
}
2.2 ODE数据的生成机制
ODE数据的生成是Stage 2训练的关键前置步骤。生成脚本位于 /Causal-Forcing/get_causal_ode_data_framewise.py,其核心逻辑是:
为什么需要ODE数据? 论文指出,直接用多步扩散模型作为教师进行蒸馏效果不佳,因为学生模型难以在少数步骤内模仿教师的多步去噪过程。ODE Distillation的思想是:先让教师模型跑完完整的采样轨迹,记录下每一步的状态,然后让学生模型学习这个轨迹。
# get_causal_ode_data_framewise.py 的核心采样逻辑
for progress_id, t in enumerate(scheduler.timesteps): # 48步采样
timestep = t * torch.ones([1, 21], device=device)
# 使用训练好的AR Diffusion模型进行去噪
f_cond, x0_pred_cond = model(latents, conditional_dict, timestep, clean_x=clean_latent)
f_uncond, x0_pred_uncond = model(latents, unconditional_dict, timestep, clean_x=clean_latent)
# CFG (Classifier-Free Guidance)
flow_pred = f_uncond + guidance_scale * (f_cond - f_uncond)
latents = scheduler.step(flow_pred, timestep, latents)
noisy_input.append(latents) # 记录每一步的latent状态
# 保存关键时间步的轨迹 [t=1000, t=750, t=500, t=250, 最终结果, GT]
noisy_inputs = noisy_inputs[:, [0, 12, 24, 36, -2, -1]]
这里的关键设计是使用了Teacher Forcing模式的采样:在每一步去噪时,模型可以看到干净的GT帧作为条件(clean_x=clean_latent)。这确保了生成的ODE轨迹是"理想情况下的轨迹",学生模型学习这样的轨迹才能获得正确的初始化。
三、模型架构深度分析
3.1 因果注意力机制的设计哲学
论文的核心创新之一是将传统的双向扩散Transformer改造为因果Transformer,使其支持自回归生成。模型定义位于/Causal-Forcing/wan/modules/causal_model.py。
因果注意力的数学表达:设视频序列为 $x_1, x_2, ..., x_T$,双向注意力允许每个位置看到所有位置,而因果注意力要求第 $i$ 帧只能看到第 $1$ 到 $i$ 帧的信息。但在视频生成中,完全严格的因果性会降低生成质量——因为即使是"未来帧"的生成,也可以受益于对整体语义的理解。
项目采用的方案是Block-wise Causal Attention:将视频帧按块划分,块内使用双向注意力,块间使用因果注意力。这平衡了生成质量与流式生成的需求。
# wan/modules/causal_model.py:511-566
@staticmethod
def _prepare_blockwise_causal_attn_mask(device, num_frames=21,
frame_seqlen=1560, num_frame_per_block=1):
"""
构造块级因果注意力掩码
每个帧块内部的token可以互相看到,但只能看到之前块的token
"""
def attention_mask(b, h, q_idx, kv_idx):
# kv_idx < ends[q_idx] 表示kv位置在当前块结束之前
return (kv_idx < ends[q_idx]) | (q_idx == kv_idx)
block_mask = create_block_mask(attention_mask, ...)
return block_mask
KV Cache机制:为了支持流式推理,模型实现了增量式的KV Cache。当生成新帧时,不需要重新计算之前帧的Key和Value,而是直接从Cache中读取。
# wan/modules/causal_model.py:197-245
# 当存在kv_cache时,进行增量更新
if kv_cache is not None:
# 将新的Key-Value追加到Cache中
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
# 只使用最近的KV进行注意力计算
x = attention(roped_query,
kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index],
kv_cache["v"][:, max(0, local_end_index - self.max_attention_size):local_end_index])
3.2 模型加载与前向传播
模型Wrapper设计:/Causal-Forcing/utils/wan_wrapper.py 封装了三个核心组件:
class WanDiffusionWrapper(torch.nn.Module):
"""扩散模型Wrapper,支持因果和双向两种模式"""
def __init__(self, model_name="Wan2.1-T2V-1.3B", is_causal=False, ...):
if is_causal:
self.model = CausalWanModel.from_pretrained(f"wan_models/{model_name}/")
else:
self.model = WanModel.from_pretrained(f"wan_models/{model_name}/")
# Flow Matching调度器
self.scheduler = FlowMatchScheduler(shift=timestep_shift, ...)
def forward(self, noisy_image_or_video, conditional_dict, timestep,
clean_x=None, # Teacher Forcing模式的关键参数
kv_cache=None, ...):
# 将噪声latent输入模型,预测flow(去噪方向)
flow_pred = self.model(noisy_image_or_video, t=timestep, context=prompt_embeds, ...)
# 将flow预测转换为x0预测(干净图像估计)
pred_x0 = self._convert_flow_pred_to_x0(flow_pred, noisy_image_or_video, timestep)
return flow_pred, pred_x0
Flow Matching vs DDPM:项目采用Flow Matching作为扩散框架,而非传统的DDPM。在Flow Matching中,噪声添加过程为 $x_t = (1-\sigma_t)x_0 + \sigma_t \epsilon$,模型预测的是从噪声到干净图像的方向向量(flow)。这种参数化方式在少步采样时更稳定。
# utils/scheduler.py:159-176
def add_noise(self, original_samples, noise, timestep):
"""Flow Matching的前向过程"""
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
sample = (1 - sigma) * original_samples + sigma * noise
return sample
def training_target(self, sample, noise, timestep):
"""训练目标:从当前状态指向噪声的方向"""
target = noise - sample
return target
四、三阶段训练流程详解
4.1 Stage 1: 自回归扩散训练
配置文件:/Causal-Forcing/configs/ar_diffusion_tf_framewise.yaml
核心问题:如何训练一个因果结构的扩散模型,使其能够根据历史帧生成当前帧?
论文的解决方案——Teacher Forcing:在训练时,条件帧使用真实的GT帧,而非模型自己生成的帧。这避免了训练时的误差累积,让模型总是学习"在正确条件下如何生成正确结果"。
训练Pipeline实现:/Causal-Forcing/pipeline/teacher_forcing_training.py
def inference_with_trajectory(self, noise, clean_image_or_video, **conditional_dict):
"""
Teacher Forcing训练的核心逻辑
"""
# 关键:使用干净的GT帧作为条件
# 模型同时看到干净帧和噪声帧,通过注意力掩码控制信息流
for index, current_timestep in enumerate(self.denoising_step_list):
_, output = self.generator(
noisy_image_or_video=noisy_input, # 噪声帧(需要去噪的)
conditional_dict=conditional_dict,
timestep=timestep,
clean_x=clean_image_or_video # 干净GT帧(作为条件)
)
注意力掩码的特殊设计:在Teacher Forcing模式下,输入同时包含干净帧和噪声帧,需要精心设计的掩码确保信息流的正确性。
# wan/modules/causal_model.py:569-654
@staticmethod
def _prepare_teacher_forcing_mask(device, num_frames=21, frame_seqlen=1560, ...):
"""
Teacher Forcing的注意力掩码设计
- 干净帧:只能看到自己和之前的干净帧(块级因果)
- 噪声帧:可以看到所有之前的干净帧 + 自己块内的噪声帧
"""
def attention_mask(b, h, q_idx, kv_idx):
# 干净帧的掩码
clean_mask = (q_idx < clean_ends) & (kv_idx < context_ends[q_idx])
# 噪声帧的掩码:可以看到之前的干净帧 + 自己块内
C1 = (kv_idx < noise_noise_ends[q_idx]) & (kv_idx >= noise_noise_starts[q_idx])
C2 = (kv_idx < noise_context_ends[q_idx]) & (kv_idx >= noise_context_starts[q_idx])
noise_mask = (q_idx >= clean_ends) & (C1 | C2)
return eye_mask | clean_mask | noise_mask
Loss计算:/Causal-Forcing/model/diffusion.py:48-130
def generator_loss(self, image_or_video_shape, conditional_dict, clean_latent, ...):
# 1. 随机采样timestep(每个块可以有不同的噪声水平)
index = self._get_timestep(0, 1000, batch_size, num_frame, num_frame_per_block)
timestep = self.scheduler.timesteps[index]
# 2. 向干净latent添加噪声
noisy_latents = self.scheduler.add_noise(clean_latent, noise, timestep)
# 3. 模型前向,预测flow
flow_pred, x0_pred = self.generator(
noisy_image_or_video=noisy_latents,
conditional_dict=conditional_dict,
timestep=timestep,
clean_x=clean_latent if self.teacher_forcing else None
)
# 4. Flow Matching Loss
training_target = noise - clean_latent # flow = ε - x0
loss = F.mse_loss(flow_pred, training_target)
训练超参数(来自配置文件):
- 学习率:2e-6
- Batch size:1(每个GPU)
- 优化器:AdamW(beta1=0.0, beta2=0.999)
- 训练步数:推荐 ≥2000步,5K-10K更佳
- 混合精度:bfloat16
4.2 Stage 2: ODE蒸馏 / Consistency蒸馏
方案一:ODE Distillation
动机:Stage 1训练的模型需要多步(如48步)采样才能生成高质量视频,无法满足实时需求。ODE Distillation的目标是让模型学会用少数步(如4步)完成生成。
核心思想:让模型学习预计算的ODE轨迹,使得在任意中间时间步,模型的预测都能与ODE轨迹上的下一状态匹配。
模型实现:/Causal-Forcing/model/ode_regression.py
def generator_loss(self, ode_latent, conditional_dict):
"""
ode_latent: [B, num_steps, F, C, H, W]
包含从噪声到干净的完整轨迹,按时间从大到小排列
"""
clean_latent = ode_latent[:, -1] # 最终的干净图像
target_latent = ode_latent[:, -2] # 轨迹上的下一状态
# 随机选择一个中间时间步
noisy_input, timestep = self._prepare_generator_input(ode_latent)
# 模型预测
_, pred_image_or_video = self.generator(
noisy_image_or_video=noisy_input,
conditional_dict=conditional_dict,
timestep=timestep,
clean_x=clean_latent # 仍然使用TF模式
)
# 回归Loss:预测应该落在轨迹的下一状态上
loss = F.mse_loss(pred_image_or_video, target_latent)
为什么使用双向模型作为ODE轨迹生成器? 论文指出,虽然最终部署的是因果模型,但ODE轨迹可以用更强的双向模型生成,因为ODE Distillation只需要学生模型的输出与轨迹匹配,不要求轨迹本身的生成方式必须是因果的。
方案二:Consistency Distillation(推荐)
动机:生成ODE配对数据需要大量计算资源(需要跑完完整的采样轨迹)。Consistency Distillation提供了一种无需预计算数据的替代方案。
核心思想:让模型在任何时间步的预测都映射到同一个最终结果(一致性)。具体来说,如果对同一个样本在时间步 $t$ 和 $t'$ 分别加噪后预测,两者的结果应该一致。
模型实现:/Causal-Forcing/model/naive_consistency.py
def generator_loss(self, conditional_dict, unconditional_dict, clean_latent, ema_model):
# 1. 随机选择时间步
t = self.scheduler.timesteps[timestep_idx]
# 2. 向GT加噪
latent_t = self.scheduler.add_noise(clean_latent, noise, timestep=t)
# 3. Teacher模型前向(使用CFG)
with torch.no_grad():
v_cond, _ = self.teacher(latent_t, conditional_dict, timestep, clean_x=clean_latent)
v_uncond, _ = self.teacher(latent_t, unconditional_dict, timestep, clean_x=clean_latent)
v_pred = v_uncond + guidance_scale * (v_cond - v_uncond)
# 计算下一时间步的状态
dt = (timestep - timestep_next) / 1000
latent_t_next = latent_t - dt * v_pred
# 4. 学生模型预测当前步
_, cm_pred_t = self.generator(latent_t, conditional_dict, timestep, clean_x=clean_latent)
# 5. EMA学生模型预测下一步(作为一致性目标)
ema_model.copy_to(self.generator_ema)
_, cm_pred_t_next = self.generator_ema(latent_t_next, conditional_dict, timestep_next, clean_x=clean_latent)
# 6. 一致性Loss
loss = F.mse_loss(cm_pred_t, cm_pred_t_next)
为什么用EMA学生模型作为目标? 使用指数移动平均(EMA)的参数可以提供更稳定的目标,避免训练过程中的震荡。这是一个常见于自监督学习的技巧。
4.3 Stage 3: Distribution Matching Distillation (DMD)
配置文件:/Causal-Forcing/configs/causal_forcing_dmd_framewise.yaml
核心问题:经过前两个阶段,模型已经能够用4步生成视频。但论文发现,直接在自生成帧上进行DMD训练仍存在条件分布不匹配问题。
论文的关键发现:在DMD阶段,可以使用双向教师模型,因为DMD只要求学生模型的最终输出分布与教师匹配,而不要求生成轨迹一致。双向模型通常比因果模型更强,因此是更好的教师。
DMD原理简述:DMD将蒸馏问题转化为对抗训练问题。学生模型作为生成器,另一个"fake score"网络作为判别器。训练目标是让学生模型的输出分布逼近真实数据分布(由教师模型定义)。
模型实现:/Causal-Forcing/model/dmd.py
class DMD(SelfForcingModel):
def __init__(self, args, device):
# 三个核心模型
self.generator = WanDiffusionWrapper(is_causal=True) # 学生(因果)
self.real_score = WanDiffusionWrapper(model_name="Wan2.1-T2V-14B", is_causal=False) # 教师(双向)
self.fake_score = WanDiffusionWrapper(model_name="Wan2.1-T2V-1.3B", is_causal=False) # 判别器
def _compute_kl_grad(self, noisy_image_or_video, estimated_clean_image_or_video,
timestep, conditional_dict, unconditional_dict):
"""
计算DMD梯度(论文公式7)
"""
# 1. Fake score:学生模型的分布估计
_, pred_fake_image = self.fake_score(noisy_image_or_video, conditional_dict, timestep)
# 2. Real score:教师模型的分布估计
_, pred_real_image_cond = self.real_score(noisy_image_or_video, conditional_dict, timestep)
_, pred_real_image_uncond = self.real_score(noisy_image_or_video, unconditional_dict, timestep)
pred_real_image = pred_real_image_cond + guidance_scale * (pred_real_image_cond - pred_real_image_uncond)
# 3. DMD梯度 = 学生预测 - 教师预测
grad = pred_fake_image - pred_real_image
# 4. 归一化(论文公式8)
normalizer = torch.abs(estimated_clean_image_or_video - pred_real_image).mean()
grad = grad / normalizer
return grad
def compute_distribution_matching_loss(self, image_or_video, conditional_dict, unconditional_dict):
"""
计算DMD Loss
"""
with torch.no_grad():
# 随机采样timestep
timestep = self._get_timestep(min_timestep, max_timestep, ...)
# 添加噪声
noise = torch.randn_like(image_or_video)
noisy_latent = self.scheduler.add_noise(image_or_video, noise, timestep)
# 计算KL梯度
grad, dmd_log_dict = self._compute_kl_grad(noisy_latent, image_or_video, timestep, ...)
# DMD Loss:让生成结果向教师分布移动
dmd_loss = 0.5 * F.mse_loss(image_or_video, (image_or_video - grad).detach())
return dmd_loss
训练策略:DMD采用生成器-判别器交替训练的策略。
# trainer/distillation.py:303-379
def train(self):
while True:
# 训练生成器(学生模型)
if self.step % self.config.dfake_gen_update_ratio == 0:
generator_log_dict = self.fwdbwd_one_step(batch, train_generator=True)
self.generator_optimizer.step()
if self.generator_ema is not None:
self.generator_ema.update(self.model.generator)
# 训练判别器(fake score)
critic_log_dict = self.fwdbwd_one_step(batch, train_generator=False)
self.critic_optimizer.step()
self.step += 1
关键超参数:
- 生成器与判别器更新比例:dfake_gen_update_ratio=5(每训练5次判别器,训练1次生成器)
- 教师模型:Wan2.1-T2V-14B(14B参数的双向模型)
- 学生模型:Wan2.1-T2V-1.3B(1.3B参数的因果模型)
- EMA衰减率:0.99
- 训练步数:Frame-wise推荐500步,Chunk-wise推荐100-200步
五、推理流程分析
5.1 推理入口与Pipeline选择
推理入口:/Causal-Forcing/inference.py
项目提供了两种推理Pipeline:
- CausalInferencePipeline:用于已蒸馏的少步模型(4步DMD)
- CausalDiffusionInferencePipeline:用于未蒸馏的多步扩散模型
5.2 核心推理流程
Pipeline实现:/Causal-Forcing/pipeline/causal_inference.py
def inference(self, noise, text_prompts, initial_latent=None):
"""
Causal Forcing的核心推理流程
"""
batch_size, num_frames, C, H, W = noise.shape
# 1. 文本编码
conditional_dict = self.text_encoder(text_prompts)
# 2. 初始化KV Cache(用于增量推理)
self._initialize_kv_cache(batch_size, dtype, device)
self._initialize_crossattn_cache(batch_size, dtype, device)
# 3. 如果是I2V,先处理初始图像
if initial_latent is not None:
# 将初始图像编码为latent,并用t=0更新KV Cache
self.generator(noisy_image_or_video=initial_latent, timestep=0,
kv_cache=self.kv_cache1, ...)
# 4. 逐块生成(流式输出)
all_num_frames = [self.num_frame_per_block] * num_blocks
if self.independent_first_frame and initial_latent is None:
all_num_frames = [1] + all_num_frames # 第一帧单独生成
for current_num_frames in all_num_frames:
noisy_input = noise[:, current_start_frame:current_start_frame + current_num_frames]
# 4.1 空间去噪循环(4步)
for current_timestep in self.denoising_step_list: # [1000, 750, 500, 250]
_, denoised_pred = self.generator(
noisy_image_or_video=noisy_input,
conditional_dict=conditional_dict,
timestep=current_timestep,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start_frame * self.frame_seq_length
)
# 添加下一步噪声(模拟多步采样)
if index < len(self.denoising_step_list) - 1:
next_timestep = self.denoising_step_list[index + 1]
noisy_input = self.scheduler.add_noise(
denoised_pred,
torch.randn_like(denoised_pred),
next_timestep
)
# 4.2 用clean latent更新KV Cache
# 关键:用去噪后的结果作为下一块的条件
self.generator(
noisy_image_or_video=denoised_pred,
timestep=self.args.context_noise, # 可以添加少量噪声
kv_cache=self.kv_cache1,
...
)
current_start_frame += current_num_frames
# 5. VAE解码
video = self.vae.decode_to_pixel(output)
return video
5.3 KV Cache的增量更新机制
设计动机:在因果注意力中,生成第 $i$ 帧时需要计算与之前所有帧的注意力。如果每次都重新计算,计算复杂度是 $O(T^2)$。通过KV Cache,可以将复杂度降低到 $O(T)$。
实现细节:/Causal-Forcing/pipeline/causal_inference.py:286-306
def _initialize_kv_cache(self, batch_size, dtype, device):
"""
为每个Transformer Block初始化KV Cache
"""
for _ in range(30): # 30个Transformer blocks
kv_cache1.append({
"k": torch.zeros([batch_size, 32760, 12, 128]), # [B, max_seq_len, heads, head_dim]
"v": torch.zeros([batch_size, 32760, 12, 128]),
"global_end_index": torch.tensor([0]), # 当前处理到的全局位置
"local_end_index": torch.tensor([0]) # Cache中的有效长度
})
Cache更新逻辑:每次生成新帧时,将其Key和Value追加到Cache末尾。
5.4 推理命令示例
# T2V (Text-to-Video)
python inference.py \
--config_path configs/causal_forcing_dmd_framewise.yaml \
--checkpoint_path checkpoints/framewise/causal_forcing.pt \
--data_path prompts/demos.txt \
--output_folder output/framewise \
--use_ema # 使用EMA参数获得更好质量
# I2V (Image-to-Video) - 仅frame-wise模型支持
python inference.py \
--config_path configs/causal_forcing_dmd_framewise.yaml \
--checkpoint_path checkpoints/framewise/causal_forcing.pt \
--data_path prompts/i2v \
--output_folder output/i2v \
--i2v \
--use_ema
六、并行训练与分布式策略
6.1 FSDP并行策略
实现位置:/Causal-Forcing/utils/distributed.py
项目使用PyTorch的FSDP (Fully Sharded Data Parallel) 进行分布式训练,支持多节点多GPU的大规模训练。
def fsdp_wrap(module, sharding_strategy="hybrid_full", mixed_precision=True, ...):
"""
FSDP包装器
sharding_strategy选择:
- "full": 完全分片,每个GPU只保存部分参数
- "hybrid_full": 混合分片,节点内分片+节点间复制(推荐)
- "hybrid_zero2": 类似ZeRO-2,只分片梯度
"""
sharding_strategy = {
"full": ShardingStrategy.FULL_SHARD,
"hybrid_full": ShardingStrategy.HYBRID_SHARD,
"hybrid_zero2": ShardingStrategy._HYBRID_SHARD_ZERO2,
}[sharding_strategy]
# 混合精度:参数用bf16,梯度累加用fp32
mixed_precision_policy = MixedPrecision(
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
buffer_dtype=torch.float32
)
# 按参数量自动分片(min_num_params=5e7,即5000万参数)
auto_wrap_policy = partial(size_based_auto_wrap_policy, min_num_params=min_num_params)
module = FSDP(module, auto_wrap_policy=auto_wrap_policy, ...)
return module
6.2 EMA在分布式环境下的实现
class EMA_FSDP:
"""
分布式环境下的EMA实现
关键:EMA参数存储在CPU上,避免GPU显存占用
"""
def __init__(self, fsdp_module, decay=0.99):
self.decay = decay
self.shadow = {} # 存储在CPU
self._init_shadow(fsdp_module)
def update(self, fsdp_module):
# θ_ema = decay * θ_ema + (1 - decay) * θ
for n, p in fsdp_module.module.named_parameters():
self.shadow[n].mul_(self.decay).add_(p.detach().float().cpu(), alpha=1 - self.decay)
def full_state_dict(self, fsdp_module):
# 保存时用EMA参数替换当前参数
for n, p in fsdp_module.module.named_parameters():
p.data.copy_(self.shadow[n].to(dtype=p.dtype, device=p.device))
return fsdp_state_dict(fsdp_module)
七、模型保存与加载
7.1 Checkpoint结构
训练过程中,模型保存在 {logdir}/checkpoint_model_{step:06d}/model.pt。
# 保存逻辑 (trainer/distillation.py:188-210)
def save(self):
generator_state_dict = fsdp_state_dict(self.model.generator)
if self.step >= self.config.ema_start_step:
state_dict = {
"generator_ema": self.generator_ema.full_state_dict(self.model.generator),
}
else:
state_dict = {
"generator": generator_state_dict,
}
torch.save(state_dict, save_path)
7.2 Checkpoint加载
# 加载逻辑 (inference.py:68-80)
if args.checkpoint_path:
state_dict = torch.load(args.checkpoint_path, map_location="cpu")
key = 'generator_ema' if args.use_ema else 'generator'
gen_sd = state_dict[key]
# 处理FSDP的key前缀
try:
pipeline.generator.load_state_dict(gen_sd)
except RuntimeError:
fixed = {}
for k, v in gen_sd.items():
if k.startswith("model._fsdp_wrapped_module."):
k = k.replace("model._fsdp_wrapped_module.", "model.", 1)
fixed[k] = v
pipeline.generator.load_state_dict(fixed, strict=False)
八、完整项目目录结构与功能
/Causal-Forcing/
│
├── configs/ # YAML配置文件
│ ├── default_config.yaml # 默认配置(会被具体配置覆盖)
│ ├── ar_diffusion_tf_*.yaml # Stage 1: AR Diffusion训练配置
│ ├── causal_ode_*.yaml # Stage 2: ODE Distillation配置
│ ├── causal_cd_*.yaml # Stage 2: Consistency Distillation配置
│ └── causal_forcing_dmd_*.yaml # Stage 3: DMD配置
│
├── model/ # 模型定义(核心算法逻辑)
│ ├── base.py # 基类:模型初始化、timestep采样
│ ├── diffusion.py # Stage 1模型:CausalDiffusion + Flow Matching
│ ├── ode_regression.py # Stage 2模型:ODE轨迹回归
│ ├── naive_consistency.py # Stage 2替代:Consistency Distillation
│ ├── dmd.py # Stage 3模型:Distribution Matching Distillation
│ ├── causvid.py # CausVid基线实现
│ └── gan.py # GAN对抗训练扩展
│
├── pipeline/ # 推理与训练Pipeline(控制流程)
│ ├── causal_inference.py # 少步推理:DMD蒸馏后的模型推理
│ ├── causal_diffusion_inference.py # 多步推理:完整扩散采样
│ ├── teacher_forcing_training.py # Teacher Forcing训练Pipeline
│ ├── self_forcing_training.py # Self-Forcing训练Pipeline(对比方法)
│ └── bidirectional_training.py # 双向模型训练Pipeline
│
├── trainer/ # 训练器(训练循环逻辑)
│ ├── diffusion.py # DiffusionTrainer: Stage 1训练循环
│ ├── ode.py # ODETrainer: Stage 2 ODE训练循环
│ ├── naive_cd.py # ConsistencyDistillationTrainer: CD训练
│ └── distillation.py # ScoreDistillationTrainer: DMD训练
│
├── utils/ # 工具函数
│ ├── dataset.py # 数据集:LMDB读取、TextDataset、I2V数据
│ ├── lmdb_.py # LMDB存储工具
│ ├── wan_wrapper.py # 模型Wrapper:Diffusion + TextEncoder + VAE
│ ├── distributed.py # 分布式:FSDP封装、EMA、进程启动
│ ├── scheduler.py # FlowMatchScheduler: 加噪、去噪、时间表
│ ├── loss.py # Loss函数:Flow、x0、Noise预测
│ └── misc.py # 杂项:seed设置等
│
├── wan/ # Wan模型核心实现(底层架构)
│ ├── modules/
│ │ ├── causal_model.py # CausalWanModel: 因果Transformer实现
│ │ ├── model.py # WanModel: 原始双向Transformer
│ │ ├── attention.py # 注意力实现
│ │ ├── vae.py # 3D Video VAE
│ │ ├── t5.py # UMT5-XXL文本编码器
│ │ └── tokenizers.py # Tokenizer
│ ├── configs/ # 模型配置(层数、维度等)
│ ├── distributed/ # 并行策略(FSDP、上下文并行)
│ └── utils/ # 采样器(DPM-Solver、UniPC)
│
├── wan_models/ # 预训练权重目录
│ ├── Wan2.1-T2V-1.3B/ # 学生模型权重
│ └── Wan2.1-T2V-14B/ # 教师模型权重(更强)
│
├── checkpoints/ # 训练产物
│ ├── framewise/ # Frame-wise模型
│ │ ├── ar_diffusion.pt # Stage 1结果
│ │ ├── causal_ode.pt # Stage 2结果
│ │ └── causal_forcing.pt # 最终模型
│ └── chunkwise/ # Chunk-wise模型
│
├── dataset/ # 数据目录
│ ├── clean_data/ # GT latent(LMDB)
│ └── ODE6KCausal_*/ # ODE轨迹数据
│
├── scripts/ # 训练脚本
├── demo_utils/ # Demo工具
├── long_video/ # 长视频生成扩展
│
├── train.py # 训练入口
├── inference.py # 推理入口
├── get_causal_ode_data_*.py # ODE数据生成脚本
└── README.md # 项目说明
九、总结
Causal Forcing项目通过三阶段渐进式蒸馏,系统性地解决了自回归视频生成中的误差累积问题:
- Stage 1 (Teacher Forcing AR Diffusion):在训练时使用GT帧作为条件,学习理想的条件分布,避免训练时的误差累积。
- Stage 2 (ODE/CD Distillation):将多步扩散压缩为少步,让学生模型学会在有限步数内准确预测去噪方向。ODE Distillation通过预计算轨迹监督,Consistency Distillation则通过自监督学习。
- Stage 3 (DMD with Bidirectional Teacher):使用更强的双向模型作为教师,通过对抗训练让学生模型的输出分布逼近教师分布。关键洞察是DMD不要求师生轨迹一致,因此可以使用更强的双向教师。
整个框架设计精妙,每个阶段都解决了特定问题,最终实现了在单张RTX 4090上的实时高质量视频生成。
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)