Diffusers 中的 Krea2Transformer2DModel:Krea 2 单流 MMDiT 流匹配 Transformer 架构全解
【免费下载链接】diffusers🤗 Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch.项目地址: https://gitcode.com/GitHub_Trending/di/diffusers
Krea 2(K2)是 Krea AI 推出的流匹配(flow-matching)文生图模型,其核心骨干是本文要剖析的Krea2Transformer2DModel——一个带分组查询注意力(GQA)的单流 MMDiT(Mixture-of-Experts of Diffusion Transformers)Transformer。本文以 Krea2Transformer2DModel 官方 API 文档 为主体,结合 transformer_krea2.py 源码 与 pipeline_krea2.py、模型测试 逐一讲解该模型的输入输出契约、模块组成、关键超参数与源码级实现原理,帮助你理解 Krea 2 在 Diffusers 生态中的落地方式,并能在本地用Krea2Transformer2DModel直接加载、推理与微调。
一、模型定位:Krea 2 的单流 MMDiT 骨干
官方文档对Krea2Transformer2DModel的定义非常精炼:它是Krea 2 所使用的单流 MMDiT 流匹配 Transformer(the single-stream MMDiT flow-matching transformer)。在 Krea 2 管线文档 中,整个模型家族被进一步描述为:
- 以Qwen3-VL 文本编码器提供条件:不取最后一层隐藏状态,而是逐 token 抽取 12 个 decoder 层的隐藏状态,堆叠后在 Transformer 内部由一个轻量text-fusion(文本融合)阶段融合;
- 图像解码使用Qwen-Image VAE(f8,16 个潜变量通道);
- 整个骨干是**单流(single-stream)**设计——文本与图像 token 拼接成一条
[text, image]序列,由同一组 Transformer block 处理。
在 Diffusers 源码中,该模型位于 src/diffusers/models/transformers/transformer_krea2.py,类定义继承自ModelMixin, ConfigMixin, AttentionMixin, PeftAdapterMixin(transformer_krea2.py#L339),因此天然支持 Diffusers 的配置序列化、注意力处理器替换、LoRA 适配器(PEFT)加载等能力,并通过 src/diffusers/models/transformers/init.py 导出为顶层 APIdiffusers.Krea2Transformer2DModel。
与其他 MMDiT(如 Flux)的差异
从源码结构看,Krea 2 的骨干与 Flux 的FluxTransformer2DModel属于同一"单流 MMDiT + RoPE + AdaGN 调制"家族(Krea2RotaryPosEmbed在源码中直接标注为从FluxPosEmbed复制修改而来,见 transformer_krea2.py#L309-L336),但有三处关键差异:
- 文本条件不是单一向量,而是一个层堆叠(layer stack):
encoder_hidden_states的 shape 是(batch, text_seq_len, num_text_layers, text_hidden_dim),需要先经过Krea2TextFusion融合。 - 注意力是 GQA + q/k RMSNorm + sigmoid 输出门控,而非 Flux 的 MQA 风格。
- 时间调制向量在所有 block 间共享:每个 block 只学习一张可加的调制表(
scale_shift_table),见下文"时间调制"小节。
二、输入 / 输出契约(forward 签名)
Krea2Transformer2DModel.forward的完整签名如下(transformer_krea2.py#L456-L466):
def forward( self, hidden_states: torch.Tensor, # (batch_size, image_seq_len, in_channels) 已打包的带噪图像潜变量 encoder_hidden_states: torch.Tensor, # (batch_size, text_seq_len, num_text_layers, text_hidden_dim) timestep: torch.Tensor, # (batch_size,) 流匹配时间,范围 [0, 1] position_ids: torch.Tensor, # (text_seq_len + image_seq_len, 3),(t, h, w) 旋转坐标 encoder_attention_mask: torch.Tensor | None = None, # (batch_size, text_seq_len) 布尔掩码 attention_kwargs: dict[str, Any] | None = None, return_dict: bool = True, ) -> Transformer2DModelOutput | tuple[torch.Tensor]各输入的含义与约束(与源码 docstring 一致):
| 参数 | Shape | 说明 |
|---|---|---|
hidden_states | (B, image_seq_len, in_channels) | patchify 打包后的带噪图像潜变量。in_channels = vae_channels * patch_size²,默认 64 = 16(Qwen-Image VAE 通道)× 2² |
encoder_hidden_states | (B, text_seq_len, num_text_layers, text_hidden_dim) | 逐 token 堆叠的文本编码器隐藏状态栈,默认num_text_layers=12 |
timestep | (B,) | 流匹配时间,1表示纯噪声、0表示干净数据;管线中由t / num_train_timesteps归一化得到(pipeline_krea2.py#L644) |
position_ids | (text_seq_len + image_seq_len, 3) | 拼接序列的(t, h, w)旋转坐标;文本行全零,图像行是潜变量网格坐标(transformer_krea2.py#L477-L479) |
encoder_attention_mask | (B, text_seq_len) | 标记有效文本 token 的布尔掩码;全部有效时传None |
attention_kwargs | dict | 含scale键时,在本次前向期间对 LoRA 适配器设置缩放系数 |
return_dict | bool | 为True时返回Transformer2DModelOutput,否则返回(velocity,)元组 |
输出是**流匹配速度(velocity)**张量,shape 为(batch_size, image_seq_len, in_channels),只对应图像 token(文本 token 在输出前被切掉,见 transformer_krea2.py#L526)。
position_ids 的形状校验
源码对position_ids有显式校验:必须是二维且最后一维为 3,否则抛出ValueError(transformer_krea2.py#L492-L493)。这是因为 RoPE 需要(t, h, w)三个坐标轴,分别对应三个axes_dims_rope维度。
三、默认配置与关键超参数
模型构造函数的全部默认值如下(transformer_krea2.py#L391-L411),@register_to_config保证这些参数会被持久化到模型配置中:
| 参数 | 默认值 | 含义 |
|---|---|---|
in_channels | 64 | patchify 后的潜变量通道数(vae_channels * patch_size²) |
num_layers | 28 | 主干 Transformer block 数量 |
attention_head_dim | 128 | 每个注意力头的维度;总隐藏维度 =head_dim * num_heads= 6144 |
num_attention_heads | 48 | 查询头数量 |
num_key_value_heads | 12 | GQA 的键/值头数量(48 / 12 = 4 组) |
intermediate_size | 16384 | 每个 block 内 SwiGLU MLP 的隐藏维度 |
timestep_embed_dim | 256 | 正弦时间嵌入在 MLP 之前的宽度 |
text_hidden_dim | 2560 | 被消费的文本编码器隐藏维度 |
num_text_layers | 12 | 每个 token 堆叠的文本编码器层数 |
text_num_attention_heads | 20 | text-fusion block 的查询头数 |
text_num_key_value_heads | 20 | text-fusion block 的键/值头数 |
text_intermediate_size | 6912 | text-fusion block 中 SwiGLU MLP 的隐藏维度 |
num_layerwise_text_blocks | 2 | 沿层轴(逐 token)应用的 text-fusion block 数 |
num_refiner_text_blocks | 2 | 沿 token 序列应用的 text-fusion block 数 |
axes_dims_rope | (32, 48, 48) | 注意力头维度在(t, h, w)三个旋转位置轴上的切分 |
rope_theta | 1000.0 | RoPE 的基频 |
norm_eps | 1e-5 | 所有 RMSNorm 的 epsilon |
源码中的硬性约束:sum(axes_dims_rope) == attention_head_dim必须成立,否则直接抛ValueError(transformer_krea2.py#L415-L418)。默认值 32+48+48=128,正好等于attention_head_dim。
四、模块组成与数据流
模型由以下子模块组成(transformer_krea2.py#L420-L454):
self.img_in = nn.Linear(in_channels, hidden_size, bias=True) # 图像 token 输入投影 self.time_embed = Krea2TimestepEmbedding(timestep_embed_dim, hidden_size) self.time_mod_proj = nn.Linear(hidden_size, 6 * hidden_size, bias=True) # 产生 6 路调制向量 self.text_fusion = Krea2TextFusion(...) # 文本层栈融合 self.txt_in = Krea2TextProjection(text_hidden_dim, hidden_size, ...) self.rotary_emb = Krea2RotaryPosEmbed(theta=rope_theta, axes_dim=list(axes_dims_rope)) self.transformer_blocks = nn.ModuleList([Krea2TransformerBlock(...) for _ in range(num_layers)]) self.final_layer = Krea2FinalLayer(hidden_size, out_channels=in_channels, eps=norm_eps)前向数据流(对应 transformer_krea2.py#L498-L531):
- 时间路径:
timestep → time_embed → GELU → time_mod_proj,得到共享调制向量temb_mod,shape 为(B, 1, 6*hidden_size)。 - 文本路径:
encoder_hidden_states → Krea2TextFusion → Krea2TextProjection(txt_in),将 4D 层栈压成 3D 文本特征序列。 - 图像路径:
hidden_states → img_in线性投影。 - 拼接:文本与图像 token 沿序列维度
cat,形成单流[text, image]序列。 - 位置编码:
rotary_emb(position_ids)计算(t, h, w)三维 RoPE。 - 主循环:28 个
Krea2TransformerBlock依次处理拼接序列(支持梯度检查点)。 - 输出:切掉文本 token,仅保留图像 token,过
final_layer输出速度。
1. Krea2TextFusion:层栈融合器
这是 Krea 2 区别于大多数扩散 Transformer 的核心设计(transformer_krea2.py#L176-L222)。输入(B, seq, num_text_layers, dim)的处理分三步:
- layerwise 阶段:reshape 为
(B*seq, num_text_layers, dim),用num_layerwise_text_blocks个Krea2TextFusionBlock(无 RoPE、无时间调制的 pre-norm block)沿层轴做自注意力——对每个 token 独立地在 12 个文本层之间交换信息; - 投影压缩:用
nn.Linear(num_text_layers, 1, bias=False)把层轴压成 1(permute 后线性层作用于层轴); - refiner 阶段:用
num_refiner_text_blocks个同样的 block沿 token 序列精修,这一步才接收attention_mask。
这种"先在层维融合、再在 token 维精修"的两段式结构,是为了把 Qwen3-VL 多个中间层的信息高效压缩成一条文本特征序列。
2. Krea2TransformerBlock:调制 + 注意力 + SwiGLU
主干 block(transformer_krea2.py#L225-L255)是标准的 pre-norm 残差结构,关键在共享调制 + 每块可加表的时间条件机制:
modulation = temb.unflatten(-1, (6, -1)) + self.scale_shift_table prescale, preshift, pregate, postscale, postshift, postgate = modulation.unbind(-2) attn_out = self.attn((1.0 + prescale) * self.norm1(hidden_states) + preshift, ...) hidden_states = hidden_states + pregate * attn_out ff_out = self.ff((1.0 + postscale) * self.norm2(hidden_states) + postshift) hidden_states = hidden_states + postgate * ff_outtemb是(B, 1, 6*hidden_size),在所有 block 间共享;每个 block 只额外学习一个scale_shift_table(6 × hidden_size的可加参数),从而用极小的参数开销把时间信息注入注意力和 FFN 两个残差支路,并带上门控系数。
3. Krea2Attention:GQA + q/k RMSNorm + sigmoid 门控
自注意力层(transformer_krea2.py#L100-L144)特点:
- GQA 投影:
to_q投影num_heads个头,to_k/to_v只投影num_kv_heads个头; - q/k 归一化:
norm_q、norm_k使用Krea2RMSNorm对 query/key 逐头做 RMSNorm(头维度上); - RoPE:
image_rotary_emb存在时对 q/k 应用旋转位置嵌入; - sigmoid 输出门:注意力输出乘上
torch.sigmoid(to_gate(hidden_states))。
其默认处理器Krea2AttnProcessor(transformer_krea2.py#L54-L97)有一个值得注意的实现细节:没有使用enable_gqa标志,而是手动repeat_interleave复制 key/value 头。源码注释说明,Krea 2 始终带文本 padding mask,而 torch 的 SDPA 内核中只有 math 和 cuDNN 同时支持 mask 与enable_gqa,且 math 内核会物化完整的[B, H, L, L]注意力矩阵;手动重复头后结果完全一致、所有内核都接受,且上下文并行路径(拒绝enable_gqa)也能继续工作。
4. Krea2RMSNorm:零中心缩放归一化
Krea2RMSNorm(transformer_krea2.py#L37-L51)实现了一个特殊约定:有效乘数是1 + weight,且 weight 初始化为全零,以匹配 Krea 2 官方 checkpoint 的格式。激活值会 upcast 到 float32 做 RMSNorm,再转回原 dtype;模型的_keep_in_fp32_modules配置保证所有 norm 权重保持 float32。
5. Krea2TimestepEmbedding:cos-first 正弦时间嵌入
时间嵌入(transformer_krea2.py#L258-L275)使用cos-first 正弦嵌入,输入时间缩放 1000 倍,再接两层 MLP(GELU-tanh 激活)。它刻意保持序列维度为 1,使每 block 的调制向量可以广播到所有 token 上。
6. Krea2FinalLayer:自适应 RMSNorm + 输出投影
输出层(transformer_krea2.py#L292-L306)使用2 × hidden_size的调制表做 scale/shift 调制,然后线性投影回in_channels得到速度。源码注释强调它被保留为单个模块并列入_no_split_modules,以便在 device-mapped 推理时让调制表、norm 与投影保持共置。
五、在 Krea 2 管线中的调用方式
Krea2Transformer2DModel由Krea2Pipeline实例化使用(pipeline_krea2.py#L172)。管线中与 Transformer 交互的关键点:
- patchify 打包:
_pack_latents将(B, C, H, W)潜变量重排为(B, H/p * W/p, C*p*p)的 token 序列(pipeline_krea2.py#L357-L363),p = patch_size = 2;去噪结束后_unpack_latents再还原。管线中的image_processor使用vae_scale_factor * patch_size作为整体缩放因子。 - 时间归一化:调度器时间步
t除以num_train_timesteps归一化到[0, 1]再喂给 Transformer(pipeline_krea2.py#L644)。 - CFG 双前向:启用分类器自由引导时,对正负两条 prompt 各调用一次 Transformer,再按 Krea 2 约定
noise_pred + guidance_scale * (noise_pred - neg_noise_pred)合并(pipeline_krea2.py#L646-L666)。 - TDM/turbo 蒸馏检查点:管线通过
is_distilled配置区分 base(midtrain)与 TDM(distilled)版本——base 建议num_inference_steps=28, guidance_scale=4.5,turbo 建议num_inference_steps=8, guidance_scale=0.0(详见 Krea 2 管线文档);蒸馏版还会使用固定的时间偏移mu=1.15(pipeline_krea2.py#L192)。
六、测试覆盖:功能、内存、torch.compile 与 LoRA
仓库为Krea2Transformer2DModel提供了完整的测试矩阵(test_models_transformer_krea2.py),可作为理解模型行为的参考:
- Krea2TransformerTesterConfig:定义微型配置(
head_dim=8, num_heads=4, num_kv_heads=2, in_channels=16, text_hidden_dim=16, num_text_layers=3, text_seq_len=4, 2×2 图像网格),并把position_ids构造成"文本行全零 + 图像行网格坐标",同时故意将最后一个文本 token 标记为 padding,以覆盖 key-padding mask 路径(test_models_transformer_krea2.py#L119-L121)。 - ModelTesterMixin:核心前向/配置测试;
- MemoryTesterMixin:显存优化测试;
- TorchCompileTesterMixin:
torch.compile兼容性测试,覆盖(4,4)/(4,8)/(8,8)三种 shape; - TrainingTesterMixin:训练与梯度检查点测试(期望的检查点模块集合为
{"Krea2Transformer2DModel"}); - AttentionTesterMixin / LoraTesterMixin:注意力处理器替换与 LoRA 适配测试。
从测试配置可以看出,encoder_hidden_states的 4D 层栈 shape、(t, h, w)三维 RoPE、文本 padding 掩码这三条是模型对外契约中最容易出错的部分,也是阅读与复用该模型时最需要留意的输入约定。
七、快速上手示例
直接实例化并调用Krea2Transformer2DModel(参考测试中的 dummy 输入构造方式):
import torch from diffusers import Krea2Transformer2DModel model = Krea2Transformer2DModel.from_pretrained("path/to/krea2-transformer") model.to("cuda").eval() batch, text_seq, img_tokens = 1, 16, 64 # 64 = 8x8 潜变量网格 dtype = torch.bfloat16 latents = torch.randn(batch, img_tokens, model.config.in_channels, device="cuda", dtype=dtype) text_stack = torch.randn(batch, text_seq, model.config.num_text_layers, model.config.text_hidden_dim, device="cuda", dtype=dtype) timestep = torch.tensor([0.5], device="cuda", dtype=dtype) # position_ids: 文本行全零,图像行填充 (t, h, w) 网格坐标 position_ids = torch.zeros(text_seq + img_tokens, 3, device="cuda") grid_h = torch.arange(8, device="cuda").repeat_interleave(8) grid_w = torch.arange(8, device="cuda").repeat(8) position_ids[text_seq:, 1] = grid_h position_ids[text_seq:, 2] = grid_w with torch.no_grad(): out = model( hidden_states=latents, encoder_hidden_states=text_stack, timestep=timestep, position_ids=position_ids, encoder_attention_mask=torch.ones(batch, text_seq, dtype=torch.bool, device="cuda"), ) print(out.sample.shape) # torch.Size([1, 64, 64])如需端到端生成,建议直接使用Krea2Pipeline(文生图、TDM/turbo 蒸馏版以及 Modular 管线的完整示例见 Krea 2 管线文档),并按 checkpoint 类型选择采样参数:Base用num_inference_steps=28, guidance_scale=4.5;TDM/Turbo用num_inference_steps=8, guidance_scale=0.0。
八、小结
Krea2Transformer2DModel是 Krea 2 单流 MMDiT 流匹配骨干的 Diffusers 实现,其技术要点可归纳为:
- 单流 MMDiT:文本与图像 token 拼接为一条序列统一处理;
- 文本层栈融合:
Krea2TextFusion先沿层轴、再沿 token 轴融合 Qwen3-VL 的 12 层隐藏状态; - GQA 注意力:
num_attention_heads / num_key_value_heads = 4组,配 q/k RMSNorm、三维(t, h, w)RoPE 与 sigmoid 输出门; - 共享时间调制:单一调制向量 + 每 block 可加调制表,贯穿注意力和 FFN 两个支路;
- 工程化细节:
_no_split_modules/_keep_in_fp32_modules/_repeated_blocks等配置支撑设备映射推理、混合精度与模型切分,测试矩阵覆盖训练、内存、torch.compile、注意力处理器与 LoRA 全链路。
理解这些设计,既能帮助你正确复用该骨干做文生图推理,也为将其改造为其他流匹配任务(如 img2img、视频生成)或进行 LoRA 微调提供了清晰的源码级参考。
【免费下载链接】diffusers🤗 Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch.项目地址: https://gitcode.com/GitHub_Trending/di/diffusers
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考