#!/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 两个槽轮流写,/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/") 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 到两个槽之一,/state 软链指向最新完整的那个。 原地覆盖唯一的 state/ 目录时,保存中途崩溃会把唯一的续训点也毁掉;这里总是写到 「不是当前 state 指向」的那个槽,写完再用 rename 原子换软链,任一时刻 state 都指向一份 完整的存档。不删任何目录,磁盘占两份(8 卡约 140 GB)。--resume /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()