JoyAI-Image MMDiT代码深潜:Transformer3DModel、RoPE与wanx调制机制,一次讲透
【免费下载链接】JoyAI-ImageJoyAI-Image is the unified multimodal foundation model for image understanding, text-to-image generation, and instruction-guided image editing.项目地址: https://gitcode.com/gh_mirrors/jo/JoyAI-Image
JoyAI-Image 是一款统一的图像多模态基础模型,核心是8B 多模态大语言模型(MLLM)+ 16B 多模态扩散 Transformer(MMDiT),支持图像理解、文生图与指令驱动图像编辑。本文带你深潜它的生成引擎——Transformer3DModel的 MMDiT 实现,讲清三块最关键的代码:双流 Transformer 结构、3D RoPE 旋转位置编码,以及 WanX 风格的调制(modulation)机制。不用数学功底,也能看懂"一张图是怎么被 Transformer 一步步生成的"。
一图看懂:MMDiT 在 JoyAI-Image 中的位置
先看整体架构:MLLM 负责"理解",MMDiT Block 负责"生成与编辑",两者通过共享的 MLLM–MMDiT 接口闭环协作。右侧放大的正是本文要拆解的 MMDiT Block——双流(图像流/文本流)、Self-Attention、Gate、Scale & Shift 全部能在代码里找到对应物。
这个架构让 MMDiT 同时承担两种任务:生成时输入纯噪声 Token,编辑时额外拼接输入图像 Token。这个"多输入图像"能力就藏在Transformer3DModel.forward里——源码见 models.py,6 维输入b n c t h w中n是多张输入图,代码把最后一张移到首位后,沿时间维t拍平拼接,用同一套 3D 卷积与 RoPE 处理。
三大核心代码文件速览
| 文件 | 职责 |
|---|---|
| models.py | Transformer3DModel主干 +MMDoubleStreamBlock双流块 |
| posemb_layers.py | 3D 网格 RoPE:频率表构建与旋转应用 |
| modulate_layers.py | WanX 调制表与modulate/apply_gate门控 |
Transformer3DModel:从 3D 卷积 Patch 到输出投影
主干类定义在 models.py,默认超参数一目了然:
hidden_size=3072、heads_num=24→ 每个注意力头 128 维patch_size=[1,2,2]:时间维不切,空间维 2×2 卷积切块mm_double_blocks_depth=20:20 层 MMDiT 双流通路rope_dim_list=[16,56,56]:RoPE 三轴维度分配(16+56+56=128,正好等于 head_dim)
前向流程(forward)可以拆成 5 步:
- Patch 化:
img_in用 3D 卷积把 latent(b,c,t,h,w)切成 token 序列并展平; - 条件嵌入:
WanTimeTextImageEmbedding把扩散时间步timestep编码成temb和 6 组调制向量vec(time_proj_dim=hidden_size*6),同时把文本特征投影到隐藏维度; - RoPE 准备:
get_rotary_pos_embed按(tt,th,tw)网格算出图像 token 的旋转频率; - 20 层双流通路:
img和txt序列每层同步更新; - 反 Patch 输出:
proj_out+unpatchify把 token 还原回 3D latent,交给 VAE 解码成图片。
💡 多 GPU 场景下,
_fsdp_shard_conditions(models.py)声明 FSDP 以MMDoubleStreamBlock为单位分片,配合--hsdp-shard-dim参数即可多卡推理。
MMDoubleStreamBlock:双流如何"合流再分流"
MMDoubleStreamBlock(models.py)是 MMDiT 的灵魂,一个 block 内发生四件事:
- 各自调制:图像流和文本流分别用自己的调制表
img_mod/txt_mod从vec取出 6 个参数(两组 shift、scale、gate),先对各自特征做Norm → 调制; - QK-Norm:Q、K 各自过
RMSNorm(models.py)再投影,稳定注意力 logits——这是近年 DiT 的常见稳定性技巧; - 拼接做联合 Self-Attention:
q = cat(img_q, txt_q)后一起进 Flash Attention(attention.py),图像 token 直接"看见"每个文字 token,这是 MMDiT 与普通 DiT 的本质区别; - 分流加残差:注意力输出按长度切回图像/文本两段,各自
apply_gate门控后加残差,再过第二次调制的 MLP。
注意一个细节:文本分支的 RoPE 被raise NotImplementedError显式禁用(models.py)——因为文本是 1D 序列,位置已由 LLM 编码体现,只对图像 token 施加空间旋转更合理。
RoPE 深潜:三轴网格如何编码"第几行第几列"
RoPE(旋转位置编码)不新增参数,而是把位置信息"旋转"进 Q/K。核心在 posemb_layers.py:
- 建网格:
get_meshgrid_nd(L14-L56)对 3D 尺寸(tt, th, tw)生成t/h/w三轴坐标网格; - 分轴算频率:
get_nd_rotary_pos_embed(L177-L268)按rope_dim_list=[16,56,56]把 128 维 head_dim 拆给时间、高、宽三轴,各轴用自己的坐标做一维 RoPE 后拼接; - 实数形式:
use_real=True直接返回(cos, sin)而非复数——因为 Flash Attention 和 TensorRT 对 complex64 支持不佳(L197-L198); - 施加旋转:
apply_rotary_emb(L142-L174)用x·cos + rotate_half(x)·sin完成二维平面旋转,rotate_half把向量按相邻两维配对,模拟复数乘法。
为什么这样设计有效?RoPE 的关键性质是"注意力分数只依赖相对位置"。图像 token 的(t,h,w)坐标编码后,模型能天然感知"左边/上边/第 3 帧"这样的空间关系——这正是 JoyAI-Image 空间编辑能力的底层支撑。
时间轴还有个小机关:get_rotary_pos_embed里target_ndim = 3, ndim = 5 - 2,2D 输入会自动在前面补1变成(1, th, tw)(models.py),为视频/多图扩展预留了统一接口。
WanX 调制机制:一张"参数表"控制整个网络
modulate_layers.py(仅 80 行)藏着 DiT 调制的精髓。WanX 风格与 SD3/Flux 不同:调制不是每层 MLP 计算,而是查一张零初始化的可学习表。
ModulateWan 里modulate_table形状为(1, 6, hidden_size),forward 只做一件事:
# 时间步嵌入 + 调制表 → 切成 6 份:shift1, scale1, gate1, shift2, scale2, gate2 return [o.squeeze(1) for o in (self.modulate_table + x).chunk(self.factor, dim=1)]随后两个小函数各司其职:
modulate:x * (1 + scale) + shift,把"当前去噪步"的信息写进特征;apply_gate:残差连接前的软门控x * gate,控制每层更新幅度。
零初始化意味着训练初期 gate≈0、shift/scale≈0,网络输出几乎不受调制影响——这是一种非常温和的"渐进式学习"设计,也让 16B 参数的 MMDiT 训练更稳。
这套设计为什么值得学
回看Transformer3DModel的完整配方,几乎每个选择都能对应到一个工程问题:
- 双流通路→ 让文本与图像在每层平等交互,编辑指令精确落地到具体区域;
- QK-Norm + RMSNorm→ 大模型注意力稳定性的"保险丝";
- 3D 分轴 RoPE→ 无参化空间定位,天然支持任意分辨率(
tt,th,tw由输入决定); - WanX 调制表→ 用极小的额外参数承载扩散时间步条件,训练稳定且推理零开销;
- 变长 Flash Attention(get_cu_seqlens)→ 批内不同文本长度各算各的,不浪费算力。
如果你想动手验证,只需两步:pip install -e .安装依赖,然后运行图像编辑脚本python inference.py --ckpt-root <路径> --image test_images/test_1.jpg --prompt "Turn the plate blue",inference_und.py则负责理解侧——详见 README.md 的 Quick Start 章节。配合 ComfyUI 节点 joyai_image_comfyui/nodes.py 与 transformer 配置,还能直观看到线上权重hidden_size=4096、40 层、patch [1,2,2]与代码默认值的差异——同一套结构,放大参数即可得到更强的模型,这也是这套 MMDiT 代码最大的价值:简洁、正交、可扩展。
【免费下载链接】JoyAI-ImageJoyAI-Image is the unified multimodal foundation model for image understanding, text-to-image generation, and instruction-guided image editing.项目地址: https://gitcode.com/gh_mirrors/jo/JoyAI-Image
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考