news 2026/9/10 9:43:00

Diffusers 中的 Krea2Transformer2DModel:Krea 2 单流 MMDiT 流匹配 Transformer 架构全解

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Diffusers 中的 Krea2Transformer2DModel:Krea 2 单流 MMDiT 流匹配 Transformer 架构全解

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),但有三处关键差异:

  1. 文本条件不是单一向量,而是一个层堆叠(layer stack)encoder_hidden_states的 shape 是(batch, text_seq_len, num_text_layers, text_hidden_dim),需要先经过Krea2TextFusion融合。
  2. 注意力是 GQA + q/k RMSNorm + sigmoid 输出门控,而非 Flux 的 MQA 风格。
  3. 时间调制向量在所有 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_kwargsdictscale键时,在本次前向期间对 LoRA 适配器设置缩放系数
return_dictboolTrue时返回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_channels64patchify 后的潜变量通道数(vae_channels * patch_size²
num_layers28主干 Transformer block 数量
attention_head_dim128每个注意力头的维度;总隐藏维度 =head_dim * num_heads= 6144
num_attention_heads48查询头数量
num_key_value_heads12GQA 的键/值头数量(48 / 12 = 4 组)
intermediate_size16384每个 block 内 SwiGLU MLP 的隐藏维度
timestep_embed_dim256正弦时间嵌入在 MLP 之前的宽度
text_hidden_dim2560被消费的文本编码器隐藏维度
num_text_layers12每个 token 堆叠的文本编码器层数
text_num_attention_heads20text-fusion block 的查询头数
text_num_key_value_heads20text-fusion block 的键/值头数
text_intermediate_size6912text-fusion block 中 SwiGLU MLP 的隐藏维度
num_layerwise_text_blocks2沿层轴(逐 token)应用的 text-fusion block 数
num_refiner_text_blocks2沿 token 序列应用的 text-fusion block 数
axes_dims_rope(32, 48, 48)注意力头维度在(t, h, w)三个旋转位置轴上的切分
rope_theta1000.0RoPE 的基频
norm_eps1e-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):

  1. 时间路径timestep → time_embed → GELU → time_mod_proj,得到共享调制向量temb_mod,shape 为(B, 1, 6*hidden_size)
  2. 文本路径encoder_hidden_states → Krea2TextFusion → Krea2TextProjection(txt_in),将 4D 层栈压成 3D 文本特征序列。
  3. 图像路径hidden_states → img_in线性投影。
  4. 拼接:文本与图像 token 沿序列维度cat,形成单流[text, image]序列。
  5. 位置编码rotary_emb(position_ids)计算(t, h, w)三维 RoPE。
  6. 主循环:28 个Krea2TransformerBlock依次处理拼接序列(支持梯度检查点)。
  7. 输出:切掉文本 token,仅保留图像 token,过final_layer输出速度。

1. Krea2TextFusion:层栈融合器

这是 Krea 2 区别于大多数扩散 Transformer 的核心设计(transformer_krea2.py#L176-L222)。输入(B, seq, num_text_layers, dim)的处理分三步:

  1. layerwise 阶段:reshape 为(B*seq, num_text_layers, dim),用num_layerwise_text_blocksKrea2TextFusionBlock(无 RoPE、无时间调制的 pre-norm block)沿层轴做自注意力——对每个 token 独立地在 12 个文本层之间交换信息;
  2. 投影压缩:用nn.Linear(num_text_layers, 1, bias=False)把层轴压成 1(permute 后线性层作用于层轴);
  3. 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_out

temb(B, 1, 6*hidden_size)在所有 block 间共享;每个 block 只额外学习一个scale_shift_table6 × 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_qnorm_k使用Krea2RMSNorm对 query/key 逐头做 RMSNorm(头维度上);
  • RoPEimage_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 管线中的调用方式

Krea2Transformer2DModelKrea2Pipeline实例化使用(pipeline_krea2.py#L172)。管线中与 Transformer 交互的关键点:

  1. 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作为整体缩放因子。
  2. 时间归一化:调度器时间步t除以num_train_timesteps归一化到[0, 1]再喂给 Transformer(pipeline_krea2.py#L644)。
  3. CFG 双前向:启用分类器自由引导时,对正负两条 prompt 各调用一次 Transformer,再按 Krea 2 约定noise_pred + guidance_scale * (noise_pred - neg_noise_pred)合并(pipeline_krea2.py#L646-L666)。
  4. 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:显存优化测试;
  • TorchCompileTesterMixintorch.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 类型选择采样参数:Basenum_inference_steps=28, guidance_scale=4.5TDM/Turbonum_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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/10 9:42:45

如何用 Social-Analyzer 一个用户名排查上千社媒平台:实战笔记

如何用 Social-Analyzer 一个用户名排查上千社媒平台:实战笔记 【免费下载链接】social-analyzer API, CLI, and Web App for analyzing and finding a persons profile in 1000 social media \ websites 项目地址: https://gitcode.com/GitHub_Trending/so/socia…

作者头像 李华
网站建设 2026/9/10 9:41:39

ZLMediaKit-windows64 启动失败与推流不通的完整排障指南

简介:本资源为最新编译的ZLMediaKit Windows 64位流媒体服务器发行版,面向音视频开发工程师、直播系统搭建者及边缘推流场景实践者,解决Windows环境下开箱即用、低延迟部署流媒体服务的核心需求。压缩包共61个文件,含核心可执行文…

作者头像 李华
网站建设 2026/9/10 9:40:40

亚马逊选品新思路:从供给端找断层,避开红海竞争

“需求大、竞争少”这种选品思路,现在基本属于正确的废话。你打开任何一篇选品教程,都会看到类似的告诫,可真到实操环节,你会发现但凡能用数据工具直接看出来的“蓝海”,早被铺货的人踏成红海了。我自己做了几年亚马逊…

作者头像 李华
网站建设 2026/9/10 9:40:16

Buzz零门槛离线语音转文字完整指南:3步把会议录音变成纪要

Buzz零门槛离线语音转文字完整指南:3步把会议录音变成纪要 【免费下载链接】buzz Buzz transcribes and translates audio offline on your personal computer. Powered by OpenAIs Whisper. 项目地址: https://gitcode.com/GitHub_Trending/buz/buzz Buzz 是…

作者头像 李华
网站建设 2026/9/10 9:38:44

高斯混合MCMC线性地震反演:从正演模型到后验分布

简介:一套面向本硕博教研人群的线性地震反演Matlab仿真资源,聚焦高斯混合马尔科夫-蒙特卡洛(GM-MCMC)算法的编程实现与原理验证。资源包共13个文件,压缩后约1.9MB,其中包含9个M脚本/函数、2个MAT数据文件、…

作者头像 李华