Instructions to use teawhite/ActionRoPE with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Wan2.2
How to use teawhite/ActionRoPE with Wan2.2:
# No code snippets available yet for this library. # To use this model, check the repository files and the library's documentation. # Want to help? PRs adding snippets are welcome at: # https://github.com/huggingface/huggingface.js
- Notebooks
- Google Colab
- Kaggle
Download code/actionrope/train.py from teawhite/ActionRoPE: direct link, hf CLI and curl.
- Browser
- Download file 38.5 kB
-
https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/actionrope/train.py
- Command line
-
hf download hf://teawhite/ActionRoPE/code/actionrope/train.py
-
curl -L -o train.py https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/actionrope/train.py
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()} | |
| # ---------------------------------------------------------------------------- | |
| # 验证 | |
| # ---------------------------------------------------------------------------- | |
| 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() | |