
Diffusers 中的 Krea2Transformer2DModelKrea 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/diffusersKrea 2K2是 Krea AI 推出的流匹配flow-matching文生图模型其核心骨干是本文要剖析的Krea2Transformer2DModel——一个带分组查询注意力GQA的单流 MMDiTMixture-of-Experts of Diffusion TransformersTransformer。本文以 Krea2Transformer2DModel 官方 API 文档 为主体结合 transformer_krea2.py 源码 与 pipeline_krea2.py、模型测试 逐一讲解该模型的输入输出契约、模块组成、关键超参数与源码级实现原理帮助你理解 Krea 2 在 Diffusers 生态中的落地方式并能在本地用Krea2Transformer2DModel直接加载、推理与微调。一、模型定位Krea 2 的单流 MMDiT 骨干官方文档对Krea2Transformer2DModel的定义非常精炼它是Krea 2 所使用的单流 MMDiT 流匹配 Transformerthe single-stream MMDiT flow-matching transformer。在 Krea 2 管线文档 中整个模型家族被进一步描述为以Qwen3-VL 文本编码器提供条件不取最后一层隐藏状态而是逐 token 抽取 12 个 decoder 层的隐藏状态堆叠后在 Transformer 内部由一个轻量text-fusion文本融合阶段融合图像解码使用Qwen-Image VAEf816 个潜变量通道整个骨干是**单流single-stream**设计——文本与图像 token 拼接成一条[text, image]序列由同一组 Transformer block 处理。在 Diffusers 源码中该模型位于 src/diffusers/models/transformers/transformer_krea2.py类定义继承自ModelMixin, ConfigMixin, AttentionMixin, PeftAdapterMixintransformer_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 stackencoder_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-L466def 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 16Qwen-Image VAE 通道× 2²encoder_hidden_states(B, text_seq_len, num_text_layers, text_hidden_dim)逐 token 堆叠的文本编码器隐藏状态栈默认num_text_layers12timestep(B,)流匹配时间1表示纯噪声、0表示干净数据管线中由t / num_train_timesteps归一化得到pipeline_krea2.py#L644position_ids(text_seq_len image_seq_len, 3)拼接序列的(t, h, w)旋转坐标文本行全零图像行是潜变量网格坐标transformer_krea2.py#L477-L479encoder_attention_mask(B, text_seq_len)标记有效文本 token 的布尔掩码全部有效时传Noneattention_kwargsdict含scale键时在本次前向期间对 LoRA 适配器设置缩放系数return_dictbool为True时返回Transformer2DModelOutput否则返回(velocity,)元组输出是**流匹配速度velocity**张量shape 为(batch_size, image_seq_len, in_channels)只对应图像 token文本 token 在输出前被切掉见 transformer_krea2.py#L526。position_ids 的形状校验源码对position_ids有显式校验必须是二维且最后一维为 3否则抛出ValueErrortransformer_krea2.py#L492-L493。这是因为 RoPE 需要(t, h, w)三个坐标轴分别对应三个axes_dims_rope维度。三、默认配置与关键超参数模型构造函数的全部默认值如下transformer_krea2.py#L391-L411register_to_config保证这些参数会被持久化到模型配置中参数默认值含义in_channels64patchify 后的潜变量通道数vae_channels * patch_size²num_layers28主干 Transformer block 数量attention_head_dim128每个注意力头的维度总隐藏维度 head_dim * num_heads 6144num_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必须成立否则直接抛ValueErrortransformer_krea2.py#L415-L418。默认值 324848128正好等于attention_head_dim。四、模块组成与数据流模型由以下子模块组成transformer_krea2.py#L420-L454self.img_in nn.Linear(in_channels, hidden_size, biasTrue) # 图像 token 输入投影 self.time_embed Krea2TimestepEmbedding(timestep_embed_dim, hidden_size) self.time_mod_proj nn.Linear(hidden_size, 6 * hidden_size, biasTrue) # 产生 6 路调制向量 self.text_fusion Krea2TextFusion(...) # 文本层栈融合 self.txt_in Krea2TextProjection(text_hidden_dim, hidden_size, ...) self.rotary_emb Krea2RotaryPosEmbed(thetarope_theta, axes_dimlist(axes_dims_rope)) self.transformer_blocks nn.ModuleList([Krea2TransformerBlock(...) for _ in range(num_layers)]) self.final_layer Krea2FinalLayer(hidden_size, out_channelsin_channels, epsnorm_eps)前向数据流对应 transformer_krea2.py#L498-L531时间路径timestep → time_embed → GELU → time_mod_proj得到共享调制向量temb_modshape 为(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, biasFalse)把层轴压成 1permute 后线性层作用于层轴refiner 阶段用num_refiner_text_blocks个同样的 block沿 token 序列精修这一步才接收attention_mask。这种先在层维融合、再在 token 维精修的两段式结构是为了把 Qwen3-VL 多个中间层的信息高效压缩成一条文本特征序列。2. Krea2TransformerBlock调制 注意力 SwiGLU主干 blocktransformer_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_table6 × hidden_size的可加参数从而用极小的参数开销把时间信息注入注意力和 FFN 两个残差支路并带上门控系数。3. Krea2AttentionGQA 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头维度上RoPEimage_rotary_emb存在时对 q/k 应用旋转位置嵌入sigmoid 输出门注意力输出乘上torch.sigmoid(to_gate(hidden_states))。其默认处理器Krea2AttnProcessortransformer_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零中心缩放归一化Krea2RMSNormtransformer_krea2.py#L37-L51实现了一个特殊约定有效乘数是1 weight且 weight 初始化为全零以匹配 Krea 2 官方 checkpoint 的格式。激活值会 upcast 到 float32 做 RMSNorm再转回原 dtype模型的_keep_in_fp32_modules配置保证所有 norm 权重保持 float32。5. Krea2TimestepEmbeddingcos-first 正弦时间嵌入时间嵌入transformer_krea2.py#L258-L275使用cos-first 正弦嵌入输入时间缩放 1000 倍再接两层 MLPGELU-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-L363p patch_size 2去噪结束后_unpack_latents再还原。管线中的image_processor使用vae_scale_factor * patch_size作为整体缩放因子。时间归一化调度器时间步t除以num_train_timesteps归一化到[0, 1]再喂给 Transformerpipeline_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配置区分 basemidtrain与 TDMdistilled版本——base 建议num_inference_steps28, guidance_scale4.5turbo 建议num_inference_steps8, guidance_scale0.0详见 Krea 2 管线文档蒸馏版还会使用固定的时间偏移mu1.15pipeline_krea2.py#L192。六、测试覆盖功能、内存、torch.compile 与 LoRA仓库为Krea2Transformer2DModel提供了完整的测试矩阵test_models_transformer_krea2.py可作为理解模型行为的参考Krea2TransformerTesterConfig定义微型配置head_dim8, num_heads4, num_kv_heads2, in_channels16, text_hidden_dim16, num_text_layers3, text_seq_len4, 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)三种 shapeTrainingTesterMixin训练与梯度检查点测试期望的检查点模块集合为{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, devicecuda, dtypedtype) text_stack torch.randn(batch, text_seq, model.config.num_text_layers, model.config.text_hidden_dim, devicecuda, dtypedtype) timestep torch.tensor([0.5], devicecuda, dtypedtype) # position_ids: 文本行全零图像行填充 (t, h, w) 网格坐标 position_ids torch.zeros(text_seq img_tokens, 3, devicecuda) grid_h torch.arange(8, devicecuda).repeat_interleave(8) grid_w torch.arange(8, devicecuda).repeat(8) position_ids[text_seq:, 1] grid_h position_ids[text_seq:, 2] grid_w with torch.no_grad(): out model( hidden_stateslatents, encoder_hidden_statestext_stack, timesteptimestep, position_idsposition_ids, encoder_attention_masktorch.ones(batch, text_seq, dtypetorch.bool, devicecuda), ) print(out.sample.shape) # torch.Size([1, 64, 64])如需端到端生成建议直接使用Krea2Pipeline文生图、TDM/turbo 蒸馏版以及 Modular 管线的完整示例见 Krea 2 管线文档并按 checkpoint 类型选择采样参数Base用num_inference_steps28, guidance_scale4.5TDM/Turbo用num_inference_steps8, guidance_scale0.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),仅供参考