惯性聚合 高效追踪和阅读你感兴趣的博客、新闻、科技资讯
阅读原文 在惯性聚合中打开

推荐订阅源

MyScale Blog
MyScale Blog
博客园 - 司徒正美
A
About on SuperTechFans
Vercel News
Vercel News
H
Hackread – Cybersecurity News, Data Breaches, AI and More
爱范儿
爱范儿
I
InfoQ
freeCodeCamp Programming Tutorials: Python, JavaScript, Git & More
博客园_首页
Google DeepMind News
Google DeepMind News
T
Tailwind CSS Blog
奇客Solidot–传递最新科技情报
奇客Solidot–传递最新科技情报
F
Fortinet All Blogs
S
SegmentFault 最新的问题
阮一峰的网络日志
阮一峰的网络日志
D
Docker
钛媒体:引领未来商业与生活新知
钛媒体:引领未来商业与生活新知
G
Google Developers Blog
Stack Overflow Blog
Stack Overflow Blog
M
MIT News - Artificial intelligence
Jina AI
Jina AI
H
Help Net Security
量子位
IT之家
IT之家

博客园 - Dsp Tian

DiT (Diffusion Transformer) 骨干网络详解 Flow Matching 原理与 MNIST 条件生成实践 Claude Code 自动推送测试 ssh端口转发 【Python】使用uv虚拟环境 解决ModuleNotFoundError: No module named 'pkg_resources' 配置Nginx反向代理 【Python】大模型工具调用 Claude Code配置Qwen3-Coder OpenCode + Oh My OpenCode配置Qwen3-Coder 【Python】vllm部署调用Qwen3-VL make指定安装目录 解决colcon编译卡死 【Python】调用C++ 深度学习(Grad-CAM) 深度学习(CVAE) 深度学习(DBBNet重参数化) 深度学习(视觉注意力SeNet/CbmaNet/SkNet/EcaNet) 深度学习(ACNet重参数化) 深度学习(RepVGG重参数化) 深度学习(修改onnx文件batchsize) 【Python】生成git仓库贡献热力图 深度学习(onnx量化) 深度学习(pytorch量化) cmake构建后执行命令
MMDiT 骨干网络详解
Dsp Tian · 2026-08-29 · via 博客园 - Dsp Tian

Multi-Modal Diffusion Transformer —— SD3(Stable Diffusion 3)等文生图大模型的核心骨干。

本文以一份最小可运行实现(MNIST + Flow Matching)逐模块拆解 MMDiT 的骨干结构,并在关键处用单流 DiT 作为对比参照,帮助理解「MMDiT 到底在 DiT 的基础上改了什么、为什么这么改」。文末附 mmdit.py 完整代码。


一、整体架构总览

MMDiT 的完整前向数据流(mmdit.py:230-267):

输入 x (B,1,28,28)                      输入 labels (B,)
      │                                       │
      │  patch_embed (Conv2d)                 │  label_embed
      ▼                                       ▼
图像 token (B, N_img, D)              标签嵌入 (B, D)
      │  + pos_embed_img                      │
      ▼                                       ├──▶ 注入文本 token
(B, N_img, D)                                 │
      │                                       ▼
      │                          文本 token (B, N_txt, D)
      │                          = text_query_tokens + label_embed + pos_embed_txt
      │                                       │
      └───────────────┬───────────────────────┘
                      ▼
               MMDiTBlock × depth
         (联合注意力 + 每模态独立 LayerNorm/MLP + adaLN-Zero)
                      │
                      ▼
              仅取图像 token → norm_final → final_linear → unpatchify
                      │
                      ▼
              输出速度场 v_t (B,1,28,28)

核心思想一句话:把「文本/条件」从 DiT 里的一条全局向量,升级为一条独立的 token 序列,与图像 token 一起走 Transformer 块——在注意力层里跨模态交互,在归一化/MLP 层各自独立。

关键超参数mmdit.py:138-149):

参数 含义
image_size 28 MNIST 图像边长
patch_size 4 每个 patch 的边长,28/4=7 → 7×7=49 个图像 token
hidden_dim (D) 256 所有 token 的统一特征维度
depth 8 Transformer 块数量
num_heads 4 联合注意力的头数
num_text_tokens 8 模拟「文本」的 token 序列长度
num_classes 10 MNIST 类别数(模拟文本内容)

二、双模态 token 的构建

2.1 图像模态:patch embedding + 位置编码

# mmdit.py:161-164, 242-244
self.patch_embed = nn.Conv2d(
    in_channels, hidden_dim,
    kernel_size=patch_size, stride=patch_size, bias=True,
)

img_tokens = self.patch_embed(x)                     # (B, D, 7, 7)
img_tokens = img_tokens.flatten(2).transpose(1, 2)   # (B, 49, D)
img_tokens = img_tokens + self.pos_embed_img          # (B, 49, D)
  • stride = kernel_size = patch_size 的卷积做 patch 化(等价于 ViT 的线性 patch embedding)。
  • 展平 + 转置后得到 N_img = (28/4)² = 49 个图像 token。
  • pos_embed_img 是可学习位置编码 (1, 49, D),广播加到每个 token 上。

2.2 文本/条件模态:query token + 标签注入 + 位置编码

这是 MMDiT 相对 DiT 新增的部分(mmdit.py:171-181, 247-252):

# 1. 可学习 query token(与标签无关的通用模板)
self.text_query_tokens = nn.Parameter(
    torch.randn(1, num_text_tokens, hidden_dim) * 0.02
)
# 2. 标签嵌入 → 加到每个 text token 上作为内容注入
self.label_embed = nn.Embedding(num_classes, hidden_dim)
# 3. 文本位置编码(可学习)
self.pos_embed_txt = nn.Parameter(
    torch.randn(1, num_text_tokens, hidden_dim) * 0.02
)
# 前向:组装文本 token
txt_tokens = self.text_query_tokens.expand(B, -1, -1)   # (B, 8, D)
y_emb = self.label_embed(labels).unsqueeze(1)           # (B, 1, D)
txt_tokens = txt_tokens + y_emb                         # 标签广播注入
txt_tokens = txt_tokens + self.pos_embed_txt            # 加文本位置编码

要点:这里用「类别标签 + 可学习 query token」来模拟文本序列。在真实 SD3 里,这 8 个 token 会被替换成 T5/CLIP 编码出的真实文本 token;骨干结构(联合注意力 + 双流归一化)完全不用改。

2.3 对比 DiT:条件从「全局向量」变成「token 序列」

DiT MMDiT
条件形式 单个全局向量 c 一条 token 序列(N_txt 个)
条件进入模型的方式 只走 adaLN adaLN + 联合注意力(token 流)
交互粒度 序列级(全局 scale/shift) token 级(逐 token attention)

三、MMDiTBlock 核心详解

MMDiTBlockmmdit.py:35-123)是 MMDiT 骨干的灵魂。一个块由四部分组成:双流 LayerNorm → 联合注意力 → 双流 MLP → adaLN-Zero 调制

3.1 双流 LayerNorm(每模态独立归一化)

# mmdit.py:48-51
self.norm1_img = nn.LayerNorm(hidden_dim, elementwise_affine=False)
self.norm1_txt = nn.LayerNorm(hidden_dim, elementwise_affine=False)
self.norm2_img = nn.LayerNorm(hidden_dim, elementwise_affine=False)
self.norm2_txt = nn.LayerNorm(hidden_dim, elementwise_affine=False)
  • 每个模态在注意力子层MLP 子层各有自己的一套 LayerNorm,共 4 个。
  • elementwise_affine=False:归一化不内置可学习的 γ/β,改由 adaLN 提供的 scale/shift 完成(见 3.4)。
  • 为什么分开:图像 patch 特征(来自卷积)和文本特征(来自 query token + embedding)分布差异大,各自归一化才能让两个模态处在各自合适的尺度上。

3.2 联合注意力(跨模态交互的灵魂)

# mmdit.py:54-56(定义,共享一个 MHA)
self.attn = nn.MultiheadAttention(
    hidden_dim, num_heads, batch_first=True
)
# mmdit.py:100-109(前向)
img_norm1 = self.norm1_img(img_tokens) * (1 + s_a_img.unsqueeze(1)) + sh_a_img.unsqueeze(1)
txt_norm1 = self.norm1_txt(txt_tokens) * (1 + s_a_txt.unsqueeze(1)) + sh_a_txt.unsqueeze(1)

joint = torch.cat([img_norm1, txt_norm1], dim=1)       # (B, N_img+N_txt, D)
attn_out = self.attn(joint, joint, joint)[0]           # 联合注意力

img_attn = attn_out[:, :N_img, :]                       # 拆回图像
txt_attn = attn_out[:, N_img:, :]                       # 拆回文本
  • 两个模态归一化后拼接成一条序列,过一个共享的 MHA。
  • 于是每个图像 patch 可以直接 attention 到任意文本 token,反之亦然——这就是 token 级跨模态交互。
  • 注意力结束后再按 N_img 切分回两个模态。

与 DiT 的本质区别:DiT 的注意力只在图像 token 内部做(自注意力),条件无法进入注意力计算;MMDiT 让条件 token 也进入注意力,实现「图-文对话」。

3.3 独立 MLP

# mmdit.py:59-68
self.mlp_img = nn.Sequential(
    nn.Linear(hidden_dim, mlp_hidden),
    nn.GELU(approximate='tanh'),
    nn.Linear(mlp_hidden, hidden_dim),
)
self.mlp_txt = nn.Sequential(...)   # 结构相同,参数独立
  • 每个模态各一个 MLP(hidden_dim → 4×hidden_dim → hidden_dim)。
  • 与 LayerNorm 同理:注意力层负责「交换信息」,MLP 层负责「各自加工」,参数独立让两个模态互不迁就。

3.4 adaLN-Zero 调制

# mmdit.py:72-78
self.adaLN_modulation = nn.Sequential(
    nn.SiLU(),
    nn.Linear(hidden_dim, 12 * hidden_dim),   # 12D = 2 模态 × 6 参数
)
nn.init.zeros_(self.adaLN_modulation[-1].weight)   # 零初始化末层
nn.init.zeros_(self.adaLN_modulation[-1].bias)
# mmdit.py:92-96(前向)
mod = self.adaLN_modulation(c)          # (B, 12D)
img_mod, txt_mod = mod.chunk(2, dim=-1) # 各 (B, 6D)

s_a_img, sh_a_img, g_a_img, s_m_img, sh_m_img, g_m_img = img_mod.chunk(6, dim=-1)
s_a_txt, sh_a_txt, g_a_txt, s_m_txt, sh_m_txt, g_m_txt = txt_mod.chunk(6, dim=-1)
  • 条件向量 c(时间 + 标签)经 SiLU + Linear 输出 12Dchunk(2) 成图像/文本各 6D
  • 每个模态的 6 个参数是两组 (scale, shift, gate),分别作用于注意力子层MLP 子层
  • adaLN-Zero:末层零初始化,使训练初期每个块的残差门 gate=0,块退化为恒等映射,有利于深层网络稳定起步。
  • 应用方式:norm(x) * (1 + scale) + shift,残差乘 gate(见 mmdit.py:100-101, 112-113, 117-121)。

3.5 块内完整数据流

img_tokens ──▶ norm1_img ──┐
txt_tokens ──▶ norm1_txt ──┴──▶ cat ──▶ attn ──▶ split
                                (联合注意力)
                                    │
              img: + g_a_img · img_attn   txt: + g_a_txt · txt_attn
                                    │
              img ──▶ norm2_img ──▶ mlp_img ──▶ + g_m_img · ...
              txt ──▶ norm2_txt ──▶ mlp_txt ──▶ + g_m_txt · ...

3.6 对比 DiTBlock

组件 DiTBlock(单流) MMDiTBlock(双流)
LayerNorm 1 套(norm1/norm2 每模态 1 套(共 4 个)
注意力 图像自注意力 拼接后联合注意力
MLP 1 个 每模态 1 个(共 2 个)
adaLN 输出 6D 12D(每模态 6D)
接口 block(x, c) block(img_tokens, txt_tokens, c)

四、条件向量与双通道注入

# mmdit.py:184-188, 255-258
self.time_embed = nn.Sequential(
    nn.Linear(hidden_dim, hidden_dim * 4),
    nn.SiLU(),
    nn.Linear(hidden_dim * 4, hidden_dim),
)

t_emb = self.time_embed(timestep_embedding(t, self.hidden_dim))
y_pool = self.label_embed(labels)          # (B, D)
c = t_emb + y_pool                         # 条件向量
  • 时间步先经正弦嵌入timestep_embeddingmmdit.py:20-29)再经 MLP,得到 t_emb
  • c = t_emb + y_pool 送入每个块的 adaLN_modulation,负责全局调制。

标签信息在 MMDiT 中走了两条通道

通道 路径 粒度
全局(adaLN) label_embed → y_pool → c → adaLN 序列级
token(联合注意力) label_embed → y_emb → 注入文本 token → 注意力 token 级

DiT 只有第一条通道;MMDiT 补上了第二条,让条件既能「全局定调」,又能「局部对齐」。


五、输出层与 unpatchify

# mmdit.py:197-200, 265-267
self.norm_final = nn.LayerNorm(hidden_dim, elementwise_affine=False)
self.final_linear = nn.Linear(
    hidden_dim, patch_size * patch_size * in_channels,
)

img_tokens = self.final_linear(self.norm_final(img_tokens))  # 仅图像 token
x = self.unpatchify(img_tokens)
  • 只取图像 token 做输出投影(文本 token 在此被丢弃,其使命已完成)。
  • final_linear 把每个 token 映射回 patch_size² × in_channels = 4×4×1 = 16 维,即一个 patch 的像素值。
  • unpatchifymmdit.py:220-228)把 49 个 patch token 拼回 28×28 图像:
x = x.reshape(-1, h, w, p, p, c)    # (B, 7, 7, 4, 4, 1)
x = x.permute(0, 5, 1, 3, 2, 4)     # (B, 1, 7, 4, 7, 4)
x = x.reshape(-1, c, h * p, w * p)  # (B, 1, 28, 28)

六、Flow Matching 训练与采样

训练目标(Rectified Flow / Flow Matching,mmdit.py:322-325):

noise = torch.randn_like(images)
t = torch.rand(images.size(0), device=device)
xt = (1 - t) * noise + t * images      # 线性插值:噪声 → 图像
vt_pred = model(xt, t, labels)          # 预测速度场
loss = F.mse_loss(vt_pred, images - noise)   # 速度 = 图像 - 噪声
  • 在噪声与图像之间做线性插值得到中间态 xt,模型预测速度场 v = 图像 - 噪声,用 MSE 拟合。
  • 与 DDPM 的「预测噪声」不同,这里是「预测速度」,是 SD3 采用的 rectified flow 范式。

采样(Euler ODE 求解器,mmdit.py:273-282):

@torch.no_grad()
def generate(label, num_samples=16, num_steps=100):
    model.eval()
    x = torch.randn(num_samples, 1, 28, 28, device=device)
    dt = 1.0 / num_steps
    for i in range(num_steps):
        t = torch.full((num_samples,), i * dt, device=device)
        x = x + model(x, t, labels) * dt   # 沿速度场逐步积分
    return (x.clamp(-1, 1) + 1) / 2
  • 从纯噪声出发,用模型预测的速度场做 100 步欧拉积分,逐步「走」向图像。
  • 最后 clamp 到 [-1,1] 再线性缩放到 [0,1] 用于保存。

七、相比 DiT 的改进总结

改进 说明
双流 + 联合注意力 条件从「全局向量」升级为「token 序列」,图像 patch 可逐 token 对齐文本,实现细粒度跨模态交互
模态独立 LayerNorm / MLP 图像、文本各自适配特征分布,只在注意力层交换信息
条件双通道注入 adaLN(全局定调)+ token 流(局部对齐)并用,表达力更强

代价(诚实地说)

  • 参数量约 1.74 倍(≈17.4M vs ≈10.0M):来自双套 MLP 与 12D adaLN。
  • 联合注意力序列长度变为 N_img + N_txt,计算量随总长度平方增长。
  • 本实现用类别标签模拟文本,优势是「演示性」的;真实收益在接入 T5/CLIP 文本编码后显现。

八、参数量估算

组件 参数量
patch_embed 4,352
pos_embed_img 12,544
text_query_tokens 2,048
pos_embed_txt 2,048
label_embed 2,560
time_embed 525,568
每块注意力(共享) 263,168
每块 MLP(×2) 1,051,136
每块 adaLN(12D) 789,504
8 个块合计 ≈ 16.83M
final_linear 4,112
总计 ≈ 17.4M

附录:mmdit.py 完整代码

import math
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision
from torchvision import transforms
from torchvision.utils import save_image
from torch.utils.data import DataLoader

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
batch_size = 256
lr = 1e-4
epochs = 100
num_classes = 10


# ============================================================
# 正弦时间嵌入
# ============================================================
def timestep_embedding(t, dim, max_period=10000):
    """Sinusoidal timestep embedding (same as DiT/DDPM)."""
    half = dim // 2
    freqs = torch.exp(
        -math.log(max_period)
        * torch.arange(half, dtype=torch.float32, device=t.device)
        / half
    )
    args = t[:, None].float() * freqs[None, :]
    return torch.cat([torch.cos(args), torch.sin(args)], dim=-1)


# ============================================================
# MMDiTBlock: 双流 Transformer 块(联合注意力 + 独立 MLP)
# ============================================================
class MMDiTBlock(nn.Module):
    """
    Multi-Modal DiT Block.
    - 每个模态拥有独立的一组 LayerNorm 和 MLP 参数
    - 注意力层将所有 token 拼接后做联合注意力,实现跨模态信息交换
    - adaLN-Zero 为每个模态生成独立的 scale/shift/gate
    """

    def __init__(self, hidden_dim, num_heads, mlp_ratio=4.0):
        super().__init__()
        mlp_hidden = int(hidden_dim * mlp_ratio)

        # ---- 每模态独立的 LayerNorm(无内置 affine,由 adaLN 提供) ----
        self.norm1_img = nn.LayerNorm(hidden_dim, elementwise_affine=False)
        self.norm1_txt = nn.LayerNorm(hidden_dim, elementwise_affine=False)
        self.norm2_img = nn.LayerNorm(hidden_dim, elementwise_affine=False)
        self.norm2_txt = nn.LayerNorm(hidden_dim, elementwise_affine=False)

        # ---- 联合注意力(所有 token 共享一个 MHA) ----
        self.attn = nn.MultiheadAttention(
            hidden_dim, num_heads, batch_first=True
        )

        # ---- 每模态独立的 MLP ----
        self.mlp_img = nn.Sequential(
            nn.Linear(hidden_dim, mlp_hidden),
            nn.GELU(approximate='tanh'),
            nn.Linear(mlp_hidden, hidden_dim),
        )
        self.mlp_txt = nn.Sequential(
            nn.Linear(hidden_dim, mlp_hidden),
            nn.GELU(approximate='tanh'),
            nn.Linear(mlp_hidden, hidden_dim),
        )

        # ---- adaLN-Zero 调制网络 ----
        # 输出: 每模态 6 个参数(attn scale/shift/gate + mlp scale/shift/gate)= 12D
        self.adaLN_modulation = nn.Sequential(
            nn.SiLU(),
            nn.Linear(hidden_dim, 12 * hidden_dim),
        )
        # 零初始化末层 —— adaLN-Zero 的关键
        nn.init.zeros_(self.adaLN_modulation[-1].weight)
        nn.init.zeros_(self.adaLN_modulation[-1].bias)

    def forward(self, img_tokens, txt_tokens, c):
        """
        Args:
            img_tokens: (B, N_img, D)  图像 patch token
            txt_tokens: (B, N_txt, D)  文本/条件 token
            c:          (B, D)         条件向量 (t + label)
        Returns:
            img_tokens, txt_tokens  (各自更新后)
        """
        N_img = img_tokens.shape[1]

        # ---- 调制参数 ----
        mod = self.adaLN_modulation(c)                          # (B, 12D)
        img_mod, txt_mod = mod.chunk(2, dim=-1)                # 各 (B, 6D)

        s_a_img, sh_a_img, g_a_img, s_m_img, sh_m_img, g_m_img = img_mod.chunk(6, dim=-1)
        s_a_txt, sh_a_txt, g_a_txt, s_m_txt, sh_m_txt, g_m_txt = txt_mod.chunk(6, dim=-1)

        # ==================== 注意力子层 ====================
        # 各模态独立归一化
        img_norm1 = self.norm1_img(img_tokens) * (1 + s_a_img.unsqueeze(1)) + sh_a_img.unsqueeze(1)
        txt_norm1 = self.norm1_txt(txt_tokens) * (1 + s_a_txt.unsqueeze(1)) + sh_a_txt.unsqueeze(1)

        # 拼接 → 联合注意力
        joint = torch.cat([img_norm1, txt_norm1], dim=1)       # (B, N_img+N_txt, D)
        attn_out = self.attn(joint, joint, joint)[0]

        # 拆分回各自模态
        img_attn = attn_out[:, :N_img, :]
        txt_attn = attn_out[:, N_img:, :]

        # 门控残差连接
        img_tokens = img_tokens + g_a_img.unsqueeze(1) * img_attn
        txt_tokens = txt_tokens + g_a_txt.unsqueeze(1) * txt_attn

        # ==================== MLP 子层 ====================
        # 各模态独立归一化 + 独立 MLP
        img_norm2 = self.norm2_img(img_tokens) * (1 + s_m_img.unsqueeze(1)) + sh_m_img.unsqueeze(1)
        txt_norm2 = self.norm2_txt(txt_tokens) * (1 + s_m_txt.unsqueeze(1)) + sh_m_txt.unsqueeze(1)

        img_tokens = img_tokens + g_m_img.unsqueeze(1) * self.mlp_img(img_norm2)
        txt_tokens = txt_tokens + g_m_txt.unsqueeze(1) * self.mlp_txt(txt_norm2)

        return img_tokens, txt_tokens


# ============================================================
# MMDiT: 多模态 Diffusion Transformer 主模型
# ============================================================
class MMDiT(nn.Module):
    """
    Multi-Modal Diffusion Transformer for flow matching on MNIST.

    两个模态:
      - 图像模态: patch embedding 后的图像 token 序列
      - 文本模态: 从类别标签构造的条件 token 序列(可学习 query token + label embedding)
    """

    def __init__(
        self,
        image_size=28,
        in_channels=1,
        patch_size=4,
        hidden_dim=256,
        depth=8,
        num_heads=4,
        mlp_ratio=4.0,
        num_classes=10,
        num_text_tokens=8,        # 条件 token 数量(模拟"文本"模态的序列长度)
    ):
        super().__init__()
        self.image_size = image_size
        self.in_channels = in_channels
        self.patch_size = patch_size
        self.hidden_dim = hidden_dim
        self.num_text_tokens = num_text_tokens

        assert image_size % patch_size == 0
        self.num_img_patches = (image_size // patch_size) ** 2

        # ---- 图像模态:Patch embedding ----
        self.patch_embed = nn.Conv2d(
            in_channels, hidden_dim,
            kernel_size=patch_size, stride=patch_size, bias=True,
        )

        # 图像位置编码(可学习)
        self.pos_embed_img = nn.Parameter(
            torch.randn(1, self.num_img_patches, hidden_dim) * 0.02
        )

        # ---- 文本模态:条件 token ----
        # 1. 可学习的 query token(与标签无关的通用的 token 模板)
        self.text_query_tokens = nn.Parameter(
            torch.randn(1, num_text_tokens, hidden_dim) * 0.02
        )
        # 2. 标签嵌入 → 加到每个 text token 上作为内容注入
        self.label_embed = nn.Embedding(num_classes, hidden_dim)
        # 3. 文本位置编码(可学习)
        self.pos_embed_txt = nn.Parameter(
            torch.randn(1, num_text_tokens, hidden_dim) * 0.02
        )

        # ---- 时间嵌入:正弦编码 → MLP ----
        self.time_embed = nn.Sequential(
            nn.Linear(hidden_dim, hidden_dim * 4),
            nn.SiLU(),
            nn.Linear(hidden_dim * 4, hidden_dim),
        )

        # ---- MMDiT 块 ----
        self.blocks = nn.ModuleList([
            MMDiTBlock(hidden_dim, num_heads, mlp_ratio)
            for _ in range(depth)
        ])

        # ---- 最终输出层(仅从图像 token 解码) ----
        self.norm_final = nn.LayerNorm(hidden_dim, elementwise_affine=False)
        self.final_linear = nn.Linear(
            hidden_dim, patch_size * patch_size * in_channels,
        )

        self._init_weights()

    def _init_weights(self):
        for module in self.modules():
            if isinstance(module, nn.Linear):
                nn.init.xavier_uniform_(module.weight)
                if module.bias is not None:
                    nn.init.zeros_(module.bias)
            elif isinstance(module, nn.Conv2d):
                nn.init.xavier_uniform_(module.weight)
                if module.bias is not None:
                    nn.init.zeros_(module.bias)
            elif isinstance(module, nn.Embedding):
                nn.init.normal_(module.weight, std=0.02)
        nn.init.normal_(self.pos_embed_img, std=0.02)
        nn.init.normal_(self.pos_embed_txt, std=0.02)
        nn.init.normal_(self.text_query_tokens, std=0.02)

    def unpatchify(self, x):
        c = self.in_channels
        p = self.patch_size
        h = w = self.image_size // p

        x = x.reshape(-1, h, w, p, p, c)       # (B, h, w, p, p, c)
        x = x.permute(0, 5, 1, 3, 2, 4)        # (B, c, h, p, w, p)
        x = x.reshape(-1, c, h * p, w * p)      # (B, c, H, W)
        return x

    def forward(self, x, t, labels):
        """
        Args:
            x:       (B, C, H, W)   噪声图像
            t:       (B,)           时间步
            labels:  (B,)           类别标签
        Returns:
            (B, C, H, W)  速度场预测 v_t
        """
        B = x.shape[0]

        # ============ 图像模态 ============
        img_tokens = self.patch_embed(x)                     # (B, D, h, w)
        img_tokens = img_tokens.flatten(2).transpose(1, 2)   # (B, N_img, D)
        img_tokens = img_tokens + self.pos_embed_img

        # ============ 文本/条件模态 ============
        # 可学习 query token 扩展至 batch
        txt_tokens = self.text_query_tokens.expand(B, -1, -1)          # (B, N_txt, D)
        # 将标签信息注入每个 text token
        y_emb = self.label_embed(labels).unsqueeze(1)                   # (B, 1, D)
        txt_tokens = txt_tokens + y_emb                                # 广播相加
        txt_tokens = txt_tokens + self.pos_embed_txt

        # ============ 条件向量 ============
        t_emb = self.time_embed(timestep_embedding(t, self.hidden_dim))
        # 条件向量 c 用于 adaLN:时间嵌入 + 标签池化
        y_pool = self.label_embed(labels)                               # (B, D)
        c = t_emb + y_pool

        # ============ MMDiT 块 ============
        for block in self.blocks:
            img_tokens, txt_tokens = block(img_tokens, txt_tokens, c)

        # ============ 输出投影(仅使用图像 token) ============
        img_tokens = self.final_linear(self.norm_final(img_tokens))
        x = self.unpatchify(img_tokens)
        return x


# ============================================================
# 采样(Euler ODE 求解器)
# ============================================================
@torch.no_grad()
def generate(label, num_samples=16, num_steps=100):
    model.eval()
    x = torch.randn(num_samples, 1, 28, 28, device=device)
    labels = torch.full((num_samples,), label, device=device, dtype=torch.long)
    dt = 1.0 / num_steps
    for i in range(num_steps):
        t = torch.full((num_samples,), i * dt, device=device)
        x = x + model(x, t, labels) * dt
    return (x.clamp(-1, 1) + 1) / 2


def sample_images(epoch):
    samples = []
    for label in range(10):
        samples.append(generate(label, num_samples=8))
    samples = torch.cat(samples, dim=0)
    save_image(samples, f'mmdit_sample_{epoch}.png', nrow=8)


# ============================================================
# 训练主循环
# ============================================================
if __name__ == '__main__':
    transform = transforms.Compose([
        transforms.ToTensor(),
        transforms.Lambda(lambda x: 2 * x - 1)
    ])
    train_dataset = torchvision.datasets.MNIST(
        root='./data', train=True, download=True, transform=transform)
    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)

    model = MMDiT(
        image_size=28, in_channels=1, patch_size=4,
        hidden_dim=256, depth=8, num_heads=4,
        mlp_ratio=4.0, num_classes=10, num_text_tokens=8,
    ).to(device)

    total_params = sum(p.numel() for p in model.parameters())
    print(f'MMDiT total parameters: {total_params:,}', flush=True)

    optimizer = torch.optim.AdamW(model.parameters(), lr=lr)

    for epoch in range(epochs):
        model.train()
        total_loss = 0
        for images, labels in train_loader:
            images, labels = images.to(device), labels.to(device)
            noise = torch.randn_like(images)
            t = torch.rand(images.size(0), device=device)
            xt = (1 - t.view(-1, 1, 1, 1)) * noise + t.view(-1, 1, 1, 1) * images
            vt_pred = model(xt, t, labels)
            loss = F.mse_loss(vt_pred, images - noise)
            optimizer.zero_grad()
            loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            optimizer.step()
            total_loss += loss.item()

        print(f'Epoch [{epoch + 1}/{epochs}], Loss: {total_loss / len(train_loader):.4f}',
              flush=True)
        sample_images(epoch + 1)