DiT(Diffusion Transformer)属于扩散 Transformer 范式,是当前文生图、图生图领域的核心架构。与 LLM 的 Seq2Seq 训练流程不同,DiT 训练具有以下特点:
现有 Trainer 基于 Seq2SeqTrainer 思路,直接迁移 DiT 会遇到:
sqrt(alpha_t)*x + sqrt(1-alpha_t)*noise
sd-vae-ft-mse
0.18215
fully_shard
┌─────────────────────────────────────────────────────────┐ │ 用户入口 │ │ DiTTrainer(config) → train() → save_checkpoint() │ ├─────────────────────────────────────────────────────────┤ │ 数据流 (VAE + GeneratorDataset) │ │ ImageFolder → VAE.encode() → .npy → GeneratorDataset │ │ → batch(latent, t, y) │ ├─────────────────────────────────────────────────────────┤ │ 训练核心 (MindSpore TrainOneStepCell) │ │ q_sample(x, t, noise) → DiT(x_t, t, y) → MSE loss │ │ → TrainOneStepCell(loss, optimizer) │ ├─────────────────────────────────────────────────────────┤ │ 模型注册 (hyper-parallel ModelSpec) │ │ hyper_parallel.models.dit.__init__ → register_spec() │ └─────────────────────────────────────────────────────────┘
# 阶段 1:本地预处理(PyTorch) vae = AutoencoderKL.from_pretrained("stabilityai/sd-vae-ft-mse") latent = vae.encode(image).latent_dist.sample().mul_(0.18215) np.save("img_000.npy", latent.numpy()) # 阶段 2:远程加载(MindSpore) def latent_generator(): for fname in os.listdir("./latents"): latent = np.load(f"./latents/{fname}") # (4, 32, 32) yield latent.astype(np.float32), np.int32(0) dataset = GeneratorDataset(source=latent_generator, column_names=["x", "y"]).batch(2)
def construct(self, x, y): # 1. 随机采样 timestep t = ops.randint(0, self.num_timesteps, (x.shape[0],)) # 2. 真实 q_sample(扩散公式) noise = ops.standard_normal(x.shape) x_t = self.diffusion.q_sample(x, t, noise) # 3. DiT 预测噪声 model_output = self.model(x_t, t, y) # 4. MSE loss(noise prediction) loss = self.loss_fn(model_output, noise) return loss
# hyper_parallel/models/dit/__init__.py def _build_dit(cfg): model_name = getattr(cfg.model, 'name', 'DiT-S/2') return DiT_models[model_name]() register_spec( "DiT-S/2", ModelSpec(name="DiT-S/2", build_model_fn=_build_dit) )
参考 examples/mindspore/llama3/fsdp_tp_example.py:
examples/mindspore/llama3/fsdp_tp_example.py
os.environ["HYPER_PARALLEL_PLATFORM"] = "mindspore" from hyper_parallel import fully_shard from hyper_parallel.platform.mindspore.autograd_compat import enable_mindspore_backward_compat # 对 DiT 的 Transformer Block 应用 fully_shard for block in model.blocks: fully_shard(block, mesh=mesh["dp"])
from dit_trainer import DiTTrainer config = { 'model_name': 'DiT-S/2', 'weights_path': None, 'lr': 1e-4, 'weight_decay': 0, 'num_timesteps': 1000, } trainer = DiTTrainer(config) result = trainer.train_step(batch) # {"loss": float}
from mindspore.dataset import GeneratorDataset def latent_generator(latent_dir="./latents"): for fname in sorted(os.listdir(latent_dir)): latent = np.load(os.path.join(latent_dir, fname)) yield latent.astype(np.float32), np.int32(0) dataset = GeneratorDataset(source=latent_generator, column_names=["x", "y"]).batch(2)
from hyper_parallel.models.spec import get_spec spec = get_spec("DiT-S/2") model = spec.build_model_fn(cfg)
【RFC】HyperParallel Trainer 新增 DiT 系列模型支持
1. 需求背景 & 价值
1.1 背景
DiT(Diffusion Transformer)属于扩散 Transformer 范式,是当前文生图、图生图领域的核心架构。与 LLM 的 Seq2Seq 训练流程不同,DiT 训练具有以下特点:
1.2 当前问题
现有 Trainer 基于 Seq2SeqTrainer 思路,直接迁移 DiT 会遇到:
1.3 核心价值
2. 功能描述
2.1 DiT 最小适配器(batch 构造、loss 计算、checkpoint)
2.2 扩散训练链路
sqrt(alpha_t)*x + sqrt(1-alpha_t)*noise2.3 VAE 编码接入
sd-vae-ft-mse等标准 VAE 的编码/解码0.18215应用2.4 并行策略验证
fully_shard(FSDP)处理 DiT 的 Transformer Block 参数3. 设计方案
3.1 整体架构
3.2 数据流设计
# 阶段 1:本地预处理(PyTorch) vae = AutoencoderKL.from_pretrained("stabilityai/sd-vae-ft-mse") latent = vae.encode(image).latent_dist.sample().mul_(0.18215) np.save("img_000.npy", latent.numpy()) # 阶段 2:远程加载(MindSpore) def latent_generator(): for fname in os.listdir("./latents"): latent = np.load(f"./latents/{fname}") # (4, 32, 32) yield latent.astype(np.float32), np.int32(0) dataset = GeneratorDataset(source=latent_generator, column_names=["x", "y"]).batch(2)3.3 训练 Step 设计
def construct(self, x, y): # 1. 随机采样 timestep t = ops.randint(0, self.num_timesteps, (x.shape[0],)) # 2. 真实 q_sample(扩散公式) noise = ops.standard_normal(x.shape) x_t = self.diffusion.q_sample(x, t, noise) # 3. DiT 预测噪声 model_output = self.model(x_t, t, y) # 4. MSE loss(noise prediction) loss = self.loss_fn(model_output, noise) return loss3.4 ModelSpec 注册
# hyper_parallel/models/dit/__init__.py def _build_dit(cfg): model_name = getattr(cfg.model, 'name', 'DiT-S/2') return DiT_models[model_name]() register_spec( "DiT-S/2", ModelSpec(name="DiT-S/2", build_model_fn=_build_dit) )3.5 MindSpore 分布式接入
参考
examples/mindspore/llama3/fsdp_tp_example.py:os.environ["HYPER_PARALLEL_PLATFORM"] = "mindspore" from hyper_parallel import fully_shard from hyper_parallel.platform.mindspore.autograd_compat import enable_mindspore_backward_compat # 对 DiT 的 Transformer Block 应用 fully_shard for block in model.blocks: fully_shard(block, mesh=mesh["dp"])4. 对外 API
4.1 DiTTrainer
from dit_trainer import DiTTrainer config = { 'model_name': 'DiT-S/2', 'weights_path': None, 'lr': 1e-4, 'weight_decay': 0, 'num_timesteps': 1000, } trainer = DiTTrainer(config) result = trainer.train_step(batch) # {"loss": float}4.2 数据生成器
from mindspore.dataset import GeneratorDataset def latent_generator(latent_dir="./latents"): for fname in sorted(os.listdir(latent_dir)): latent = np.load(os.path.join(latent_dir, fname)) yield latent.astype(np.float32), np.int32(0) dataset = GeneratorDataset(source=latent_generator, column_names=["x", "y"]).batch(2)4.3 ModelSpec 获取
from hyper_parallel.models.spec import get_spec spec = get_spec("DiT-S/2") model = spec.build_model_fn(cfg)5. 使用约束
6. 测试设计
6.1 单元测试
6.2 精度对齐测试
6.3 回归测试
7. 规格 & 约束
7.1 规格
7.2 约束
8. 参考
附录:当前进度