ActionRoPE / code /actionrope /train.py
teawhite's picture
add docs+code
880dff9 verified
Raw History Blame Contribute Delete
38.5 kB
#!/usr/bin/env python3
"""ActionRoPE / plain 及四个 baseline 动作臂的全参数微调(只训 DiT + 臂的新增参数,latent 与文本都预计算)。
# 单进程(smoke,不用 DeepSpeed;AdamW 直接更新 bf16 参数、无 fp32 主权重,数字不代表训练效果)
CUDA_VISIBLE_DEVICES=0 .venv/bin/python actionrope/train.py --arm arope --limit 64 --max_steps 8 --output outputs/smoke
# 多卡:accelerate + DeepSpeed ZeRO-2(配置见 configs/accelerate_zero2*.yaml)
accelerate launch --config_file configs/accelerate_zero2.yaml actionrope/train.py --arm arope ...
为什么自己写循环而不用 DiffSynth 的 launch_training_task
--------------------------------------------------------
DiffSynth 的 runner 把 loss 藏在 pipeline 的 model_fn 里,而这里的前向是
actionrope/arope.py 的 arope_forward(世界坐标 RoPE、逐 cell 文本、mask 通道),loss 还要按
known/new 加权、帧 0 剔除;再加上定步验证、只导出 DiT 权重这些需求,套 runner 反而要绕。
loss 的形式与 DiffSynth 的 FlowMatchSFTLoss 一致(同一个 FlowMatchScheduler("Wan")、同样的
add_noise / training_target / training_weight),只是多了权重图。
各臂的差别(全在 AropeTrainModule.forward 里)
--------------------------------------------
* arope:offset_px 进 RoPE;known/new 权重(new × new_weight);可选 mask 通道;文本剥掉动作从句。
* plain:offset_px=None(普通 RoPE)、无 mask、权重全 1、文本带动作从句(text_table_eybx_mirror)。
* linear / xattn / adaln(baseline/SPEC.md):普通 RoPE、无 mask、权重全 1、文本剥掉动作从句,
动作经 `baseline.build_arm(name, dit)` 建的 ActionArm 钩子进模型(action_inputs = dataset 给的
offset_tok / delta_tok / action_idx);臂的新增参数与 DiT 同一个优化器、同 lr。
* prompt:= plain 的 ActionArm 写法(PromptArm 零参数零钩子,只核对 context 是逐 cell 的),前向与 plain 逐位一致。
所有臂都走逐 cell 文本 cross-attn,所以对照实验里唯一的变量就是"动作以什么方式进模型"。
ActionArm 必须是 AropeTrainModule 的子模块(`self.arm`):DeepSpeed 引擎只管 prepare 进去的那个 module,
臂的参数挂在外面就不会被 ZeRO 分片、也拿不到梯度。导出时 DiT 键去掉 `dit.` 前缀,臂的键保留 `arm.` 前缀
写进同一个 safetensors(metadata 记 arm 名 / arm_kwargs / n_new_params),infer.py 按键集自动装臂。
学习率调度不交给 accelerate.prepare:accelerate 包装后的 scheduler 每步会 step num_processes 次
(它假设 scheduler 是按样本数建的),warmup 会被悄悄压缩;这里手动 step,一步就是一步。
随机数:set_seed(seed, device_specific=True) 让每个 rank 的 torch / CUDA / python RNG 都不同
—— 噪声 ε(randn_like)与场景 dropout 抽签都各 rank 独立;数据分片不受影响,accelerate 会把
sampler 的 generator 从 rank 0 广播同步。不加 device_specific 时 8 张卡每步抽到同一份 ε。
"""
from __future__ import annotations
import argparse
import csv
import json
import math
import os
import platform
import socket
import sys
import time
import warnings
import zlib
os.environ.setdefault("DIFFSYNTH_SKIP_DOWNLOAD", "True")
import torch
import torch.nn as nn
from accelerate import Accelerator
from accelerate.utils import DistributedType, set_seed
from safetensors.torch import load_file, save_file
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, ROOT)
from actionrope.arope import MASK_EMBEDDING_NAME, PX_PER_LATENT, arope_forward, install_arope, known_mask, loss_weight_map # noqa: E402
from actionrope.dataset import DEFAULT_SIDECAR, TEXT_TABLES, AropeLatentDataset, arope_collate, load_sidecar # noqa: E402
from baseline import ARM_NAMES, ARM_PREFIX, arm_text_mode, build_arm, parse_arm_kwargs, split_arm_keys # noqa: E402
from diffsynth.diffusion.flow_match import FlowMatchScheduler # noqa: E402
MODEL_DIR = os.path.join(ROOT, "models/Wan2.2-TI2V-5B")
DIT_FILES = [os.path.join(MODEL_DIR, f"diffusion_pytorch_model-0000{i}-of-00003.safetensors") for i in (1, 2, 3)]
LAT_H, LAT_W = 30, 52
# 导出的 safetensors 只含 DiT,键名与原模型一致:包装 module 里 DiT 挂在 .dit 下,导出时把这个前缀去掉
DIT_PREFIX = "dit."
# ----------------------------------------------------------------------------
# 参数
# ----------------------------------------------------------------------------
def parse_args(argv=None):
ap = argparse.ArgumentParser(description="ActionRoPE 训练")
ap.add_argument("--arm", choices=list(ARM_NAMES), required=True,
help="arope(本项目)/ plain(= prompt 旧名)/ linear / xattn / prompt / adaln(baseline/SPEC.md)")
ap.add_argument("--arm_kwargs", default=None,
help="baseline 臂的构造参数 JSON,例如 xattn 复刻旧报告口径:'{\"enable_mouse\": false, \"window_frames\": 1}'")
ap.add_argument("--mask_channel", dest="mask_channel", action="store_true", default=True,
help="arope 臂是否加 known/new 掩码输入通道(零初始化 conv,默认开);plain 臂强制关")
ap.add_argument("--no_mask_channel", "--no-mask_channel", dest="mask_channel", action="store_false",
help="严格零参数版(关掉 mask 通道)")
ap.add_argument("--new_weight", type=float, default=2.0, help="new 区 loss 权重(arope 臂);--new_weight_shape dist 时是最大值")
ap.add_argument("--new_weight_shape", choices=["flat", "dist"], default="flat",
help="flat:new 区一刀切 ×new_weight;dist:按离首帧足迹边界的距离从 1 线性爬到 new_weight")
ap.add_argument("--new_weight_ramp", type=float, default=4.0, help="dist 形状的爬升距离(latent 格,16 px/格)")
ap.add_argument("--new_weight_sigma", action=argparse.BooleanOptionalAction, default=False,
help="new 区额外权重再乘 (0.5+σ):高噪声步多给、低噪声步少给")
ap.add_argument("--text_mode", choices=["scene", "scene_action"], default=None,
help="默认 arope/linear/xattn/adaln→scene(剥动作从句)、plain/prompt→scene_action(动作从句翻回真实方向)")
ap.add_argument("--text_table", default=None, help="默认按 text_mode 选 data/latent/text_table_*.pt")
ap.add_argument("--scene_dropout", type=float, default=0.1)
ap.add_argument("--mirror_labels", action=argparse.BooleanOptionalAction, default=True,
help="plain 臂:把数据集镜像的动作词翻回真实屏幕方向(用 text_table_eybx_mirror.pt);"
"--no-mirror_labels 走旧约定(原表、原串)")
ap.add_argument("--first_frame_cond", action=argparse.BooleanOptionalAction, default=True)
ap.add_argument("--train_dir", default=os.path.join(ROOT, "data/latent/train_eybx"))
ap.add_argument("--val_dir", default=os.path.join(ROOT, "data/latent/val_eybx"))
ap.add_argument("--sidecar", default=DEFAULT_SIDECAR)
ap.add_argument("--model_dir", default=MODEL_DIR)
ap.add_argument("--limit", type=int, default=None, help="smoke 用:只取这么多训练 clip")
# 世界降速(rt_ratio)闸:训练剔掉真跑慢的(<0.8,画面顿挫明显),0.8–0.9 只慢一两成、坐标又是准的,留着;
# val / 评测更严(≥0.9),否则"模型不听话"和"世界本来就慢"分不开。sidecar 的 valid 规则已放宽到 0.5 s。
ap.add_argument("--min_rt_ratio", type=float, default=0.8,
help="训练集按 meta 的 rt_ratio 剔除世界降速 clip(0 ⇒ 不剔)")
ap.add_argument("--val_min_rt_ratio", type=float, default=0.9,
help="val 集的 rt_ratio 下限(评测口径,比训练严)")
ap.add_argument("--lr", type=float, default=1e-5)
ap.add_argument("--weight_decay", type=float, default=0.01)
ap.add_argument("--warmup_steps", type=int, default=0, help="线性 warmup 步数(0 ⇒ 常数 lr)")
ap.add_argument("--grad_clip", type=float, default=1.0,
help="梯度裁剪阈值;DeepSpeed 下写进 deepspeed_config.gradient_clipping(覆盖 yaml),"
"日志里的 grad_norm 是裁剪前的全局范数")
ap.add_argument("--max_steps", type=int, default=40000)
ap.add_argument("--batch_size", type=int, default=1, help="每卡 micro batch")
ap.add_argument("--grad_accum", type=int, default=1,
help="梯度累积步数;step 计的是优化器更新数。DeepSpeed 下必须与 accelerate yaml 的 "
"gradient_accumulation_steps 一致(configs/accelerate_zero2_ga2.yaml 是 2)")
ap.add_argument("--num_workers", type=int, default=2)
ap.add_argument("--gradient_checkpointing", action=argparse.BooleanOptionalAction, default=True)
ap.add_argument("--save_every", type=int, default=1000)
ap.add_argument("--save_state", action="store_true",
help="每次导出权重时同时 accelerator.save_state(含优化器分片,8 卡约 70 GB);"
"state_A / state_B 两个槽轮流写,<output>/state 是指向最新完整槽的软链")
ap.add_argument("--val_every", type=int, default=500, help="0 ⇒ 不验证")
ap.add_argument("--val_n", type=int, default=16)
ap.add_argument("--val_timesteps", default="200,500,800")
ap.add_argument("--log_every", type=int, default=1)
ap.add_argument("--output", required=True, help="outputs/<run>")
ap.add_argument("--resume", default=None,
help="*.safetensors ⇒ 只热启动权重(step 从 0 起);目录 ⇒ accelerator.load_state(续训)")
ap.add_argument("--seed", type=int, default=0)
args = ap.parse_args(argv)
if args.text_mode is None:
args.text_mode = arm_text_mode(args.arm)
if args.text_table is None:
# plain 臂不翻动作词时用原表(镜像串不全的那张)
key = "scene_action_raw" if (args.text_mode == "scene_action" and not args.mirror_labels) else args.text_mode
args.text_table = TEXT_TABLES[key]
if args.arm != "arope":
args.mask_channel = False # 掩码通道是 ARoPE 的一部分,其它臂一律普通 RoPE + 无 mask
args.arm_kwargs = parse_arm_kwargs(args.arm_kwargs)
args.val_timesteps = [float(t) for t in args.val_timesteps.split(",") if t.strip()]
if not os.path.isabs(args.output):
args.output = os.path.join(ROOT, args.output)
return args
# ----------------------------------------------------------------------------
# 模型
# ----------------------------------------------------------------------------
def load_dit(model_dir: str, device: str = "cpu"):
"""只加载 DiT(不挂 T5 / VAE),bf16。先落在 CPU,交给 accelerator.prepare 搬上卡。"""
from diffsynth.pipelines.wan_video import ModelConfig, WanVideoPipeline
files = [os.path.join(model_dir, f"diffusion_pytorch_model-0000{i}-of-00003.safetensors") for i in (1, 2, 3)]
pipe = WanVideoPipeline.from_pretrained(
torch_dtype=torch.bfloat16, device=device,
model_configs=[ModelConfig(path=files)],
tokenizer_config=None, redirect_common_files=False,
)
return pipe.dit
def load_weights_into_dit(dit: nn.Module, path: str, mask_channel: bool, arm: nn.Module | None = None):
"""从导出的 safetensors 热启动。DiT 键必须与 dit.state_dict() 完全一致;
唯一放行的缺失是 arope 臂从 plain ckpt 起步时缺的 arope_mask_embedding.*(零初始化即可)。
`arm.` 前缀的键给动作臂(strict);ckpt 没有 arm 键而本次训练有臂 ⇒ 臂保持零初始化(从 plain / arope ckpt 起步);
ckpt 有 arm 键而本次没有臂 / 臂不同 ⇒ 报错,不能静默丢权重。"""
sd = load_file(path)
dit_sd, arm_sd = split_arm_keys(sd)
missing, unexpected = dit.load_state_dict(dit_sd, strict=False)
allowed_missing = {f"{MASK_EMBEDDING_NAME}.weight", f"{MASK_EMBEDDING_NAME}.bias"} if mask_channel else set()
bad_missing = [k for k in missing if k not in allowed_missing]
if bad_missing or unexpected:
raise RuntimeError(f"{path} 与模型键集不符:缺 {bad_missing[:5]}…({len(bad_missing)}),多 {list(unexpected)[:5]}…({len(unexpected)})")
info = {"missing": list(missing), "unexpected": list(unexpected), "n_tensors": len(sd), "n_arm_tensors": len(arm_sd)}
if arm_sd:
if arm is None:
raise RuntimeError(f"{path} 含 {len(arm_sd)} 个 {ARM_PREFIX}* 键,但本次训练的臂没有参数(--arm 与 ckpt 不符)")
arm.load_state_dict(arm_sd, strict=True) # 键集 / 形状不符(例如 xattn 的 window_frames 不同)在这里炸
info["arm_loaded"] = True
elif arm is not None and any(True for _ in arm.parameters()):
info["arm_loaded"] = False # 从无臂的 ckpt 起步:臂保持零初始化
return info
class AropeTrainModule(nn.Module):
"""把 DiT + 扩散 loss 包成一个 module,forward 直接返回 loss —— DeepSpeed 引擎要求前向在 module 内。"""
def __init__(self, dit, arm: str, mask_channel: bool, new_weight: float, first_frame_cond: bool,
gradient_checkpointing: bool, new_weight_shape: str = "flat", new_weight_ramp: float = 4.0,
new_weight_sigma: bool = False, arm_module: nn.Module | None = None):
super().__init__()
self.dit = dit
self.arm_name = arm
# baseline 动作臂(ActionArm)挂成子模块:参数进同一个优化器 / DeepSpeed 引擎,state_dict 键自带 "arm." 前缀
self.arm = arm_module
self.mask_channel = mask_channel
self.new_weight = float(new_weight)
self.new_weight_shape = new_weight_shape
self.new_weight_ramp = float(new_weight_ramp)
self.new_weight_sigma = bool(new_weight_sigma)
self.first_frame_cond = first_frame_cond
self.gradient_checkpointing = gradient_checkpointing
self.scheduler = FlowMatchScheduler("Wan")
self.scheduler.set_timesteps(1000, training=True)
def timestep_id_for(self, t: float) -> int:
return int(torch.argmin((self.scheduler.timesteps - t).abs()))
def forward(self, batch: dict, timestep_id: int, noise: torch.Tensor | None = None, train: bool = True):
x0 = batch["input_latents"] # [b, 48, 21, 30, 52] bf16
b, c = x0.shape[:2]
offset_px = batch["offset_px"] # [b, 21, 2] float32
timestep = self.scheduler.timesteps[timestep_id].view(1) # float32 cpu,与 DiffSynth 的抽法一致
if noise is None:
noise = torch.randn_like(x0)
# x_t = (1−σ) x0 + σ ε;目标 ε − x0
x_t = self.scheduler.add_noise(x0, noise, timestep)
target = self.scheduler.training_target(x0, noise, timestep)
# TI2V 首帧条件:帧 0 用干净 x0 替换,模型侧对帧 0 的 t 置 0(first_frame_cond)
x_t[:, :, 0:1] = x0[:, :, 0:1]
if self.arm_name == "arope":
known = known_mask(offset_px / PX_PER_LATENT, LAT_H, LAT_W).to(x0.device) # bool [b, 21, 30, 52]
weight = loss_weight_map(offset_px.to(x0.device), LAT_H, LAT_W, self.new_weight,
ramp_latent=self.new_weight_ramp if self.new_weight_shape == "dist" else 0.0,
sigma=(timestep.float() / 1000.0) if self.new_weight_sigma else None) # [b, 1, 21, 30, 52]
mask_input = known.to(torch.float32).unsqueeze(1) if self.mask_channel else None
offset_arg = offset_px
else:
# plain / baseline 臂:普通 RoPE、无 mask、权重全 1,动作(若有)只经 ActionArm 的钩子进模型
weight = torch.ones(b, 1, x0.shape[2], LAT_H, LAT_W, device=x0.device, dtype=torch.float32)
mask_input, offset_arg = None, None
action_inputs = None
if self.arm is not None:
# baseline/SPEC.md 的 action_inputs:dataset 已给齐四个键,原样打包(臂自己挑用哪几个)
action_inputs = {k: batch[k] for k in ("offset_px", "offset_tok", "delta_tok", "action_idx") if k in batch}
pred = arope_forward(
self.dit, x_t, timestep.to(device=x0.device, dtype=x0.dtype), batch["context"],
offset_px=offset_arg, first_frame_cond=self.first_frame_cond, mask_input=mask_input,
use_gradient_checkpointing=self.gradient_checkpointing and train,
context_ids=batch["context_ids"], arm=self.arm, action_inputs=action_inputs,
)
# 帧 0 是条件帧,不计 loss
pred, target, weight = pred[:, :, 1:], target[:, :, 1:], weight[:, :, 1:]
se = (pred.float() - target.float()) ** 2
# Σ w·(pred−target)² / Σ w;w 在通道维广播,分母乘通道数才是"加权平均"
loss_w = (se * weight).sum() / (weight.sum() * c)
mse = se.mean() # 不加权的普通 MSE,两臂可直接比
loss = loss_w * self.scheduler.training_weight(timestep).to(loss_w.device)
stats = {
"loss_w": loss_w.detach(), "mse": mse.detach(),
"new_frac": (weight > 1).float().mean().detach(),
"timestep": float(timestep.item()),
}
return loss, stats
# ----------------------------------------------------------------------------
# 导出 / 日志
# ----------------------------------------------------------------------------
def export_dit_weights(accelerator: Accelerator, model, path: str, metadata: dict) -> int:
"""取完整权重(ZeRO-2 下参数本就不分片),主进程去掉 'dit.' 前缀写 safetensors;动作臂的键保留 'arm.' 前缀
一起写进去(infer.py / load_weights_into_dit 按前缀拆)。返回文件字节数(主进程)。"""
accelerator.wait_for_everyone()
state_dict = accelerator.get_state_dict(model)
size = 0
if accelerator.is_main_process:
out = {}
for k, v in state_dict.items():
if k.startswith(DIT_PREFIX):
out[k[len(DIT_PREFIX):]] = v.detach().to("cpu").contiguous()
elif k.startswith(ARM_PREFIX):
out[k] = v.detach().to("cpu").contiguous()
# scheduler 等非 DiT / 非臂成员没有参数,这里只是保险
os.makedirs(os.path.dirname(path), exist_ok=True)
save_file(out, path, metadata={k: str(v) for k, v in metadata.items()})
os.chmod(path, 0o644) # safetensors 默认按 umask 写成 0600,推理脚本可能由别的账号跑
size = os.path.getsize(path)
print(f"[save] {path} {len(out)} 个张量 {size / 1e9:.2f} GB", flush=True)
del state_dict
accelerator.wait_for_everyone()
return size
STATE_SLOTS = ("state_A", "state_B")
def save_state_rotating(accelerator: Accelerator, output: str, trainer_state: dict):
"""accelerator.save_state 到两个槽之一,<output>/state 软链指向最新完整的那个。
原地覆盖唯一的 state/ 目录时,保存中途崩溃会把唯一的续训点也毁掉;这里总是写到
「不是当前 state 指向」的那个槽,写完再用 rename 原子换软链,任一时刻 state 都指向一份
完整的存档。不删任何目录,磁盘占两份(8 卡约 140 GB)。--resume <output>/state 照常。"""
link = os.path.join(output, "state")
cur = os.path.basename(os.path.realpath(link)) if os.path.islink(link) else None
slot = STATE_SLOTS[1] if cur == STATE_SLOTS[0] else STATE_SLOTS[0]
state_dir = os.path.join(output, slot)
accelerator.save_state(state_dir)
if accelerator.is_main_process:
with open(os.path.join(state_dir, "trainer_state.json"), "w", encoding="utf-8") as fh:
json.dump(trainer_state, fh)
tmp = link + ".tmp"
if os.path.lexists(tmp):
os.remove(tmp)
os.symlink(slot, tmp)
os.replace(tmp, link) # 只替换软链本身,不动任何目录
print(f"[save] accelerator.save_state → {state_dir}({link} → {slot})", flush=True)
accelerator.wait_for_everyone()
class CsvLog:
def __init__(self, path: str, fields: list[str]):
exists = os.path.isfile(path) and os.path.getsize(path) > 0
self.f = open(path, "a", encoding="utf-8", newline="")
self.w = csv.DictWriter(self.f, fieldnames=fields, lineterminator="\n", extrasaction="ignore")
if not exists:
self.w.writeheader()
def write(self, row: dict):
self.w.writerow(row)
self.f.flush()
def close(self):
self.f.close()
def to_device(batch: dict, device) -> dict:
return {k: (v.to(device, non_blocking=True) if isinstance(v, torch.Tensor) else v) for k, v in batch.items()}
# ----------------------------------------------------------------------------
# 验证
# ----------------------------------------------------------------------------
@torch.no_grad()
def run_validation(accelerator: Accelerator, model, module: AropeTrainModule, val_ds, tids: list[int], seed: int):
"""固定 clip 集 × 固定 timestep 集 × 固定噪声,各进程分摊 clip,主进程汇总。
ZeRO-2 的前向没有集合通信,所以各进程 clip 数不等也没关系;最后的 reduce 是集合操作,人人都要到。
返回 {t: {"loss_w", "mse", "n"}}(所有进程都拿到同样的值)。"""
model.eval()
device = accelerator.device
rank, world = accelerator.process_index, accelerator.num_processes
n_t = len(tids)
acc = torch.zeros(n_t, 3, device=device, dtype=torch.float64) # loss_w 和, mse 和, 计数
for i in range(rank, len(val_ds), world):
sample = val_ds[i]
batch = to_device(arope_collate([sample]), device)
for j, tid in enumerate(tids):
g = torch.Generator(device=device)
g.manual_seed((zlib.crc32(sample["clip_id"].encode()) * 1009 + tid * 7 + seed) % (2 ** 31))
noise = torch.randn(batch["input_latents"].shape, generator=g, device=device, dtype=batch["input_latents"].dtype)
_, stats = model(batch, tid, noise=noise, train=False)
acc[j, 0] += stats["loss_w"].double()
acc[j, 1] += stats["mse"].double()
acc[j, 2] += 1
acc = accelerator.reduce(acc, reduction="sum")
model.train()
out = {}
for j, tid in enumerate(tids):
n = max(acc[j, 2].item(), 1.0)
out[float(module.scheduler.timesteps[tid])] = {"loss_w": acc[j, 0].item() / n, "mse": acc[j, 1].item() / n, "n": int(acc[j, 2].item())}
return out
# ----------------------------------------------------------------------------
# 主流程
# ----------------------------------------------------------------------------
def main(argv=None):
args = parse_args(argv)
accelerator = Accelerator(gradient_accumulation_steps=args.grad_accum)
is_main = accelerator.is_main_process
use_ds = accelerator.distributed_type == DistributedType.DEEPSPEED
set_seed(args.seed, device_specific=True)
if use_ds:
# accelerate 的 clip_grad_norm_ 在 DeepSpeed 下只取回范数不裁剪,裁剪由引擎按 config 做;
# 让 --grad_clip 成为唯一入口,yaml 里的 gradient_clipping 只是默认值
accelerator.state.deepspeed_plugin.deepspeed_config["gradient_clipping"] = float(args.grad_clip)
ga = accelerator.state.deepspeed_plugin.deepspeed_config.get("gradient_accumulation_steps")
if ga not in ("auto", args.grad_accum):
raise SystemExit(f"--grad_accum={args.grad_accum} 与 accelerate yaml 的 gradient_accumulation_steps={ga} 不一致,"
f"请换 configs/accelerate_zero2_ga{args.grad_accum}.yaml")
accelerator.state.deepspeed_plugin.deepspeed_config["gradient_accumulation_steps"] = args.grad_accum
if is_main:
os.makedirs(args.output, exist_ok=True)
print(f"[env] world={accelerator.num_processes} distributed={accelerator.distributed_type} "
f"mixed_precision={accelerator.mixed_precision} device={accelerator.device}", flush=True)
# ---- 数据 ----
sidecar = load_sidecar(args.sidecar)
# 转场 clip 的 src_a / src_b 若在 val_dir 就剔掉(val 的帧会原样出现在转场 clip 里)
train_ds = AropeLatentDataset(args.train_dir, sidecar=sidecar, text_table=args.text_table, text_mode=args.text_mode, mirror_labels=args.mirror_labels,
scene_dropout=args.scene_dropout, limit=args.limit, seed=args.seed, verbose=is_main,
exclude_src_in=args.val_dir, min_rt_ratio=args.min_rt_ratio)
val_ds = None
if args.val_every > 0 and args.val_n > 0:
# 验证不做场景 dropout;同一张文本表,别再读一遍
val_ds = AropeLatentDataset(args.val_dir, sidecar=sidecar, text_table=train_ds.text_table, text_mode=args.text_mode, mirror_labels=args.mirror_labels,
scene_dropout=0.0, limit=args.val_n, seed=args.seed, verbose=is_main,
min_rt_ratio=args.val_min_rt_ratio)
dataloader = torch.utils.data.DataLoader(
train_ds, batch_size=args.batch_size, shuffle=True, drop_last=True, collate_fn=arope_collate,
num_workers=args.num_workers, pin_memory=True, persistent_workers=args.num_workers > 0,
)
# ---- 模型 ----
t0 = time.time()
dit = load_dit(args.model_dir, device="cpu")
install_arope(dit, mask_channel=args.mask_channel)
# baseline 臂:在 CPU 上按 dit 建参数(build_arm 会 .to(dit) —— 此时 dit 还在 CPU,随 prepare 一起上卡)
arm = build_arm(args.arm, dit, args.arm_kwargs)
n_new_params = arm.n_new_params() if arm is not None else 0
if arm is not None and not arm.zero_init_check():
raise RuntimeError(f"臂 {args.arm} 的 zero_init_check() 不通过:装上瞬间输出会偏离原模型")
resume_info = None
if args.resume and args.resume.endswith(".safetensors"):
resume_info = load_weights_into_dit(dit, args.resume, args.mask_channel, arm=arm)
if is_main:
print(f"[resume] 热启动权重 {args.resume}: {resume_info}", flush=True)
dit.train().requires_grad_(True)
if arm is not None:
arm.train().requires_grad_(True)
module = AropeTrainModule(dit, args.arm, args.mask_channel, args.new_weight, args.first_frame_cond, args.gradient_checkpointing,
new_weight_shape=args.new_weight_shape, new_weight_ramp=args.new_weight_ramp,
new_weight_sigma=args.new_weight_sigma, arm_module=arm)
n_params = sum(p.numel() for p in dit.parameters())
if is_main:
print(f"[model] DiT 加载 {time.time() - t0:.0f}s,{n_params:,} 参数,mask_channel={args.mask_channel},"
f"has {MASK_EMBEDDING_NAME}={hasattr(dit, MASK_EMBEDDING_NAME)},arm={args.arm}"
+ (f"({type(arm).__name__},新增 {n_new_params:,} 参数,arm_kwargs={args.arm_kwargs})" if arm is not None else "(无 ActionArm)"),
flush=True)
optimizer = torch.optim.AdamW(module.parameters(), lr=args.lr, weight_decay=args.weight_decay)
warm = max(int(args.warmup_steps), 0)
lr_scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lambda s: min(1.0, (s + 1) / warm) if warm > 0 else 1.0)
model, optimizer, dataloader = accelerator.prepare(module, optimizer, dataloader)
device = accelerator.device
if is_main:
clip_eff = getattr(getattr(model, "optimizer", None), "clip_grad", None) if use_ds else args.grad_clip
print(f"[env] grad_clip={args.grad_clip}(生效值 {clip_eff})", flush=True)
# RoPE 表搬上卡:plain 臂每步要查 [f*h*w,1,64] 的 complex128 表,留在 CPU 每步都要拷
module.dit.freqs = tuple(t.to(device) for t in module.dit.freqs)
start_step = 0
if args.resume and not args.resume.endswith(".safetensors"):
accelerator.load_state(args.resume)
with open(os.path.join(args.resume, "trainer_state.json"), encoding="utf-8") as fh:
start_step = int(json.load(fh)["step"])
with warnings.catch_warnings(): # 优化器还没 step 就快进 scheduler,PyTorch 会提醒一句,这里是有意的
warnings.simplefilter("ignore")
for _ in range(start_step):
lr_scheduler.step()
if is_main:
print(f"[resume] accelerator.load_state({args.resume}),从 step {start_step} 续训", flush=True)
# ---- config.json:复现所需的一切 ----
if is_main:
meta = sidecar["meta"]
cfg = {
"args": vars(args),
"sidecar_meta": {k: meta[k] for k in ("version", "gain_correction", "walk_speed_px_s", "M_px832", "n_clips", "n_valid", "max_sample_dt", "notes") if k in meta},
"dataset": {"train_total": train_ds.n_total, "train_missing": train_ds.n_missing, "train_invalid": train_ds.n_invalid,
"train_leak_excluded": train_ds.n_leak, "train_low_rt_excluded": train_ds.n_low_rt,
"train_used": len(train_ds), "text_table_keys": len(train_ds.text_table),
"val_used": len(val_ds) if val_ds is not None else 0,
"val_clip_ids": val_ds.clip_ids() if val_ds is not None else []},
"model": {"dit_files": DIT_FILES, "n_params": n_params, "mask_embedding": hasattr(dit, MASK_EMBEDDING_NAME),
"arm": args.arm, "arm_class": type(arm).__name__ if arm is not None else None,
"arm_kwargs": args.arm_kwargs, "n_new_params": n_new_params,
"arm_param_keys": [k for k, _ in arm.named_parameters()] if arm is not None else [],
"resume_info": resume_info, "start_step": start_step},
"env": {"world_size": accelerator.num_processes, "distributed_type": str(accelerator.distributed_type),
"mixed_precision": accelerator.mixed_precision, "torch": torch.__version__,
"cuda": torch.version.cuda, "gpu": torch.cuda.get_device_name(0) if torch.cuda.is_available() else None,
"python": sys.executable, "platform": platform.platform(), "hostname": socket.gethostname(),
"argv": sys.argv, "started": time.strftime("%Y-%m-%d %H:%M:%S"),
"cuda_visible_devices": os.environ.get("CUDA_VISIBLE_DEVICES")},
"deepspeed_config": accelerator.state.deepspeed_plugin.deepspeed_config if use_ds else None,
}
try:
import accelerate, deepspeed, diffsynth # noqa: F401
cfg["env"]["accelerate"] = accelerate.__version__
cfg["env"]["deepspeed"] = deepspeed.__version__
except Exception:
pass
with open(os.path.join(args.output, "config.json"), "w", encoding="utf-8") as fh:
json.dump(cfg, fh, indent=2, ensure_ascii=False, default=str)
# ---- 日志 ----
train_log = val_log = tb = None
if is_main:
train_log = CsvLog(os.path.join(args.output, "train_log.csv"),
["step", "loss", "loss_w", "mse", "new_frac", "timestep", "lr", "grad_norm", "step_time", "mem_gb", "epoch"])
val_log = CsvLog(os.path.join(args.output, "val_log.csv"), ["step", "timestep", "loss_w", "mse", "n"])
from torch.utils.tensorboard import SummaryWriter
tb = SummaryWriter(os.path.join(args.output, "tensorboard"))
# 每个进程各抽各的 timestep 与噪声(与 DiffSynth 一致),种子按 rank 错开,避免 8 张卡同一步抽到同一个 t
cpu_gen = torch.Generator().manual_seed(args.seed * 1000 + accelerator.process_index + start_step)
val_tids = [module.timestep_id_for(t) for t in args.val_timesteps]
def do_validation(step):
if val_ds is None:
return
tv = time.time()
res = run_validation(accelerator, model, module, val_ds, val_tids, args.seed)
if is_main:
for t, r in res.items():
val_log.write({"step": step, "timestep": t, **r})
tb.add_scalar(f"val/loss_w_t{int(round(t))}", r["loss_w"], step)
tb.add_scalar(f"val/mse_t{int(round(t))}", r["mse"], step)
mean_w = sum(r["loss_w"] for r in res.values()) / len(res)
mean_mse = sum(r["mse"] for r in res.values()) / len(res)
tb.add_scalar("val/loss_w_mean", mean_w, step)
tb.add_scalar("val/mse_mean", mean_mse, step)
print(f"[val] step {step} " + " ".join(f"t={t:.0f}: w={r['loss_w']:.4f} mse={r['mse']:.4f} (n={r['n']})" for t, r in res.items())
+ f" | mean w={mean_w:.4f} mse={mean_mse:.4f} {time.time() - tv:.0f}s", flush=True)
def do_save(step):
path = os.path.join(args.output, f"step-{step}.safetensors")
export_dit_weights(accelerator, model, path, {"arm": args.arm, "step": step, "mask_channel": args.mask_channel,
"new_weight": args.new_weight, "new_weight_shape": args.new_weight_shape,
"new_weight_ramp": args.new_weight_ramp, "new_weight_sigma": args.new_weight_sigma,
"text_mode": args.text_mode,
"mirror_labels": args.mirror_labels,
"arm_kwargs": json.dumps(args.arm_kwargs), "n_new_params": n_new_params})
if args.save_state:
save_state_rotating(accelerator, args.output, {"step": step, "arm": args.arm})
# ---- 训练循环 ----
model.train()
step, epoch = start_step, 0
last_saved = start_step
torch.cuda.reset_peak_memory_stats(device)
if is_main:
print(f"[train] {len(train_ds)} 个 clip,每进程 batch {args.batch_size},{accelerator.num_processes} 进程,"
f"从 step {start_step} 到 {args.max_steps}", flush=True)
micro_time = 0.0
while step < args.max_steps:
for batch in dataloader:
if step >= args.max_steps:
break
t_step = time.time()
batch = to_device(batch, device)
tid = int(torch.randint(0, 1000, (1,), generator=cpu_gen))
lr_now = lr_scheduler.get_last_lr()[0] # 本步用的 lr(scheduler.step 之后拿到的是下一步的)
# 梯度累积:accumulate() 决定这个 micro batch 是否到了同步/更新边界(sync_gradients);
# DeepSpeed 引擎按 config 自己累积并在边界 step,accelerate 的 scheduler 包装也只在边界推进
with accelerator.accumulate(model):
loss, stats = model(batch, tid)
accelerator.backward(loss) # DeepSpeed 下这一步已包含 engine.step()(裁剪 + 优化器更新)
# 单进程:真的裁剪;DeepSpeed:clip 已在 engine.step 里按 config 做过,这里只取回全局范数
grad_norm = accelerator.clip_grad_norm_(model.parameters(), args.grad_clip) if accelerator.sync_gradients else None
optimizer.step()
if accelerator.sync_gradients:
lr_scheduler.step() # DeepSpeed 的 scheduler 包装不认累积边界,每调一次就推进一次,只能在这里守门
optimizer.zero_grad(set_to_none=True)
torch.cuda.synchronize(device)
micro_time += time.time() - t_step
if not accelerator.sync_gradients:
continue # 还在累积,不算一步
step += 1
step_time, micro_time = micro_time, 0.0
loss_mean = accelerator.reduce(loss.detach().float(), reduction="mean").item()
if is_main and (step % args.log_every == 0 or step == args.max_steps):
gn = float(grad_norm) if grad_norm is not None and math.isfinite(float(grad_norm)) else float("nan")
row = {"step": step, "loss": loss_mean, "loss_w": stats["loss_w"].item(), "mse": stats["mse"].item(),
"new_frac": stats["new_frac"].item(), "timestep": stats["timestep"],
"lr": lr_now, "grad_norm": gn, "step_time": step_time,
"mem_gb": torch.cuda.max_memory_allocated(device) / 1024 ** 3, "epoch": epoch}
train_log.write(row)
for k in ("loss", "loss_w", "mse", "lr", "grad_norm", "step_time", "mem_gb", "new_frac"):
tb.add_scalar(f"train/{k}", row[k], step)
print(f"step {step}/{args.max_steps} loss {loss_mean:.4f} (w {row['loss_w']:.4f} mse {row['mse']:.4f} "
f"t {row['timestep']:.0f} new {row['new_frac']:.2f}) lr {row['lr']:.2e} gnorm {gn:.3f} "
f"{step_time:.2f}s mem {row['mem_gb']:.1f}GB", flush=True)
if args.val_every > 0 and step % args.val_every == 0:
do_validation(step)
if args.save_every > 0 and step % args.save_every == 0:
do_save(step)
last_saved = step
epoch += 1
if last_saved != step:
do_save(step)
peak = torch.tensor(torch.cuda.max_memory_allocated(device) / 1024 ** 3, device=device)
peak_max = accelerator.reduce(peak, reduction="max").item()
if is_main:
print(f"[done] step {step},本进程峰值显存 {peak.item():.1f} GB,各进程最大 {peak_max:.1f} GB", flush=True)
with open(os.path.join(args.output, "done.json"), "w", encoding="utf-8") as fh:
json.dump({"step": step, "peak_mem_gb_main": peak.item(), "peak_mem_gb_max": peak_max,
"finished": time.strftime("%Y-%m-%d %H:%M:%S")}, fh, indent=2)
train_log.close()
val_log.close()
tb.close()
accelerator.wait_for_everyone()
accelerator.end_training()
if __name__ == "__main__":
main()