Krea-2 图像生成模型实战指南:在 DiffSynth-Studio 中完成推理、低显存部署与全量/LoRA 训练
【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-Studio
Krea-2 是 Krea 团队开发的图像生成模型,本指南以 DiffSynth-Studio 仓库中的 Krea-2 官方文档 为主体,完整讲解 Krea-2-Raw 与 Krea-2-Turbo 两个版本在 DiffSynth-Studio 中的安装、快速推理、低显存部署、全量微调与 LoRA 训练全流程。读完本文,你将掌握Krea2Pipeline的加载与调用方式、全部推理与训练参数的含义与默认值、显存管理配置的底层原理,并能够直接复用仓库中提供的示例脚本完成从数据集准备到模型验证的完整工作流。
1. Krea-2 与 DiffSynth-Studio
Krea-2 是 Krea 团队发布的图像生成模型,DiffSynth-Studio 为其提供了完整的推理与训练支持。从源码结构看,Krea-2 的推理由 diffsynth/pipelines/krea2.py 中的Krea2Pipeline承载,模型整体由三部分组成:
- 文本编码器:Qwen3-VL-4B-Instruct(多模态大语言模型),负责将提示词编码为多层的 hidden states,对应源码 diffsynth/models/krea2_text_encoder.py;
- 去噪网络 DiT:SingleStreamDiT,单流混合模态 DiT,对应源码 diffsynth/models/krea2_dit.py;
- 图像 VAE:Qwen-Image VAE,负责图像与潜变量的编解码,对应源码 diffsynth/models/qwen_image_vae.py。
在Krea2Pipeline.__init__中(krea2.py),框架将调度器指定为FlowMatchScheduler("Krea-2"),即 Krea-2 采用流匹配(Flow Matching)去噪范式;同时将height_division_factor与width_division_factor均设为 16,这意味着生成图像的宽高必须是 16 的倍数。流水线由多个PipelineUnit单元按固定顺序执行(详见第 5.4 节),并通过model_fn_krea2完成单步去噪。
2. 环境安装
在使用 DiffSynth-Studio 进行 Krea-2 推理与训练之前,需要先安装 DiffSynth-Studio:
git clone https://github.com/modelscope/DiffSynth-Studio.git cd DiffSynth-Studio pip install -e .安装完成后,还需确保环境中具备torch(建议使用支持 CUDA 的版本)以及transformers、tqdm、einops等依赖。更多关于安装的细节,请参考安装依赖。
3. 快速开始:加载 Krea-2-Raw 并完成首次推理
运行以下代码可以快速加载krea/Krea-2-Raw模型并完成推理。该示例开启了显存管理,框架会自动根据剩余显存控制模型参数的加载,最低 24G 显存即可运行:
from diffsynth.pipelines.krea2 import Krea2Pipeline, ModelConfig import torch vram_config = { "offload_dtype": "disk", "offload_device": "disk", "onload_dtype": torch.float8_e4m3fn, "onload_device": "cpu", "preparing_dtype": torch.float8_e4m3fn, "preparing_device": "cuda", "computation_dtype": torch.bfloat16, "computation_device": "cuda", } pipe = Krea2Pipeline.from_pretrained( torch_dtype=torch.bfloat16, device="cuda", model_configs=[ ModelConfig(model_id="krea/Krea-2-Raw", origin_file_pattern="raw.safetensors", **vram_config), ModelConfig(model_id="Qwen/Qwen3-VL-4B-Instruct", origin_file_pattern="*.safetensors", **vram_config), ModelConfig(model_id="Qwen/Qwen-Image", origin_file_pattern="vae/diffusion_pytorch_model.safetensors", **vram_config), ], tokenizer_config=ModelConfig(model_id="Qwen/Qwen3-VL-4B-Instruct", origin_file_pattern=""), vram_limit=torch.cuda.mem_get_info("cuda")[1] / (1024 ** 3) - 1, ) prompt = "A cat standing on a stone." image = pipe(prompt, seed=0, num_inference_steps=52, cfg_scale=4.5) image.save("image.jpg")结合Krea2Pipeline.from_pretrained的实现(krea2.py),这段代码的加载过程可以做如下拆解:
- 三个
ModelConfig分别对应 Krea-2 模型的三个组件:raw.safetensors(DiT 权重)、Qwen3-VL-4B-Instruct 的全部 safetensors(文本编码器)、Qwen-Image 的 VAE 权重。origin_file_pattern用于在远端模型仓库中精确匹配需要下载的权重文件,避免下载无关文件; tokenizer_config单独指定了 Qwen3-VL-4B-Instruct 的 tokenizer,from_pretrained内部会先download_if_necessary()再通过AutoTokenizer.from_pretrained(..., max_length=512)实例化;- 加载完成后,
text_encoder、dit、vae会从模型池中按名称取出(fetch_model("krea2_text_encoder")等),pipe.vram_management_enabled会根据配置自动判断显存管理是否生效。
vram_config中的六个键对应显存管理的六个生命周期阶段,其中offload_*控制参数卸载后的存储格式与设备(此处为磁盘上的disk类型,最大限度释放显存),onload_*控制参数加载回内存时的格式(FP8 可减半内存占用),preparing_*控制参数送入 GPU 前的准备阶段(FP8),computation_*则指定实际参与前向计算的精度与设备(bfloat16 的 CUDA)。而vram_limit通过torch.cuda.mem_get_info("cuda")[1]获取 GPU 总显存并减去 1GB 作为可用显存预算,框架将据此决定哪些参数留在 GPU、哪些参数被卸载。显存管理的完整原理可参考显存管理。
4. 模型总览:Raw 与 Turbo 的完整资源索引
Krea-2 系列在 DiffSynth-Studio 中提供两个模型版本,每个版本都配套了推理、低显存推理、全量训练、训练后验证、LoRA 训练与 LoRA 验证共六类脚本:
| 模型 ID | 推理 | 低显存推理 | 全量训练 | 全量训练后验证 | LoRA 训练 | LoRA 训练后验证 | |-|-|-|-|-|-|-| | krea/Krea-2-Raw | code | code | code | code | code | code | | krea/Krea-2-Turbo | code | code | code | code | code | code |
其中Raw是基础版本,通常需要 52 步推理以获得高质量结果;Turbo是蒸馏加速版本,仅需 8 步即可出图(详见第 5.3 节)。两个版本共享相同的文本编码器与 VAE,区别仅在于 DiT 权重文件(raw.safetensors与turbo.safetensors)以及推理参数。
5. 模型推理详解
5.1 加载模型
模型统一通过Krea2Pipeline.from_pretrained加载,其完整签名(krea2.py)为:
| 参数 | 默认值 | 说明 | |-|-|-| |torch_dtype|torch.bfloat16| 模型参数与计算的默认精度 | |device| 自动检测 | 运行设备,通常为"cuda"| |model_configs|[]| 各模型组件的ModelConfig列表,指定模型 ID 与权重文件匹配模式 | |tokenizer_config|ModelConfig(model_id="Qwen/Qwen3-VL-4B-Instruct", origin_file_pattern="")| tokenizer 配置,默认自动使用 Qwen3-VL-4B-Instruct | |vram_limit|None| 显存管理预算;为None时表示不启用显存管理 |
更多关于加载模型机制的说明可参考加载模型。
5.2 推理输入参数
Krea2Pipeline.__call__的输入参数定义在 krea2.py,与官方文档完全对应:
| 参数 | 默认值 | 说明 | |-|-|-| |prompt|""| 正向提示词,描述要生成的图像内容 | |negative_prompt|""| 负向提示词,描述图像中不应该出现的内容 | |cfg_scale|3.5| Classifier-free guidance 的引导强度 | |height|1024| 图像高度,需为 16 的倍数 | |width|1024| 图像宽度,需为 16 的倍数 | |seed|None| 随机种子,None表示完全随机 | |rand_device|"cpu"| 生成随机高斯噪声矩阵的计算设备 | |num_inference_steps|52| 推理(去噪)步数 | |mu|None| 时间步动态位移(dynamic shift)参数 | |progress_bar_cmd|tqdm.tqdm| 进度条实现,可设置为lambda x: x屏蔽进度条 |
高度与宽度由Krea2Unit_ShapeChecker单元在运行时通过pipe.check_resize_height_width校验并自动对齐到合法尺寸(krea2.py)。
5.3 Turbo 版本的高效推理
Turbo 模型针对少步数蒸馏优化,仓库在 examples/krea2/model_inference/Krea-2-Turbo.py 中给出了其推荐参数,其中num_inference_steps=8、cfg_scale=1、mu=1.15是固定参数,不应随意改动:
from diffsynth.pipelines.krea2 import Krea2Pipeline, ModelConfig import torch pipe = Krea2Pipeline.from_pretrained( torch_dtype=torch.bfloat16, device="cuda", model_configs=[ ModelConfig(model_id="krea/Krea-2-Turbo", origin_file_pattern="turbo.safetensors"), ModelConfig(model_id="Qwen/Qwen3-VL-4B-Instruct", origin_file_pattern="*.safetensors"), ModelConfig(model_id="Qwen/Qwen-Image", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"), ], tokenizer_config=ModelConfig(model_id="Qwen/Qwen3-VL-4B-Instruct", origin_file_pattern=""), ) prompt = "Portrait of a woman in a blue dress, underwater, surrounded by colorful bubbles." image = pipe( prompt, seed=0, height=2048, width=2048, # The following parameters are fixed. num_inference_steps=8, cfg_scale=1, mu=1.15, ) image.save("image.jpg")注意:Turbo 的cfg_scale=1意味着不使用 CFG 引导(正向与负向分支权重相同),mu=1.15则通过流匹配调度器的动态位移机制调整时间步分布。该参数会在self.scheduler.set_timesteps(num_inference_steps, denoising_strength=1.0, dynamic_shift_len=(height // 16) * (width // 16), mu=mu)中被消费(krea2.py),其中dynamic_shift_len与图像分辨率成正比,即分辨率越高时间步位移越明显。
5.4 推理流程的内部机制(源码级)
从源码结构看,Krea2Pipeline.__call__的执行流程可以划分为三个阶段:
阶段一:预处理单元链。推理依次执行五个PipelineUnit(krea2.py):
Krea2Unit_ShapeChecker:校验并规范化宽高;Krea2Unit_NoiseInitializer:以(1, 16, height//8, width//8)的形状在rand_device上生成高斯噪声——16 是潜变量通道数,8 是 VAE 下采样倍数;Krea2Unit_PromptEmbedder:将提示词送入文本编码器。值得注意的是,编码时会在提示词前后拼接固定的系统模板"<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n"与"<|im_end|>\n<|im_start|>assistant\n",并抽取第 2、5、8、11、14、17、20、23、26、29、32、35 共 12 层的 hidden states 堆叠作为最终的条件嵌入,max_length为 512(krea2.py);Krea2Unit_InputImageEmbedder:当传入input_image(图生图场景)时用 VAE 编码输入图像,并将初始噪声按首个时间步叠加到输入潜变量上;否则直接用纯噪声初始化;Krea2Unit_PromptEmbPreCompute:将文本嵌入通过 DiT 的txtfusion与txtmlp预计算为融合后的条件,避免在每个去噪步重复计算——这由context_pre_compute=True控制。
阶段二:循环去噪。在progress_bar_cmd(self.scheduler.timesteps)的迭代中,每个时间步通过cfg_guided_model_fn调用model_fn_krea2计算噪声预测,再由self.step(self.scheduler, ...)执行调度器步进更新潜变量(krea2.py)。model_fn_krea2内部会先把潜变量按 DiT 的patch大小切分为序列(_krea2_prepare同时构建图像位置编码imgpos与掩码imgmask,与文本 token 拼接后送入 DiT),输出的 patch 序列再被重排还原为图像形状(krea2.py)。
阶段三:VAE 解码。去噪完成后加载 VAE,将最终潜变量解码为图像并保存。
6. 低显存推理
如果显存不足,请开启显存管理。仓库在 examples/krea2/model_inference_low_vram/ 中为每个模型提供了推荐的低显存配置,核心差异在于向每个ModelConfig注入vram_config并设置vram_limit(见第 3 节代码)。其工作方式为:参数以 FP8 格式在 CPU 与磁盘间流转,仅在需要参与计算时才以 bfloat16 精度加载到 CUDA,从而将显存占用压缩到最低,官方文档标注的最低可运行显存为 24G。低显存推理的代码与普通推理唯一区别就是多出的vram_config字典与vram_limit参数,推理调用方式完全一致。
7. 模型训练
7.1 训练脚本与数据集
Krea-2 系列模型统一通过 examples/krea2/model_training/train.py 训练。该脚本基于accelerate启动,内部定义了Krea2ImageTrainingModule(继承DiffusionTrainingModule),使用UnifiedDataset加载数据,并按--task分发到不同的启动器(sft:data_process走数据预处理、sft/sft:train走训练循环,train.py)。
仓库构建了样例数据集方便测试,通过以下命令即可下载:
modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "krea2/*" --local_dir ./data/diffsynth_example_dataset下载后,Krea-2-Raw 与 Krea-2-Turbo 的训练数据分别位于data/diffsynth_example_dataset/krea2/Krea-2-Raw与data/diffsynth_example_dataset/krea2/Krea-2-Turbo目录。
7.2 通用训练参数详解
train.py通过add_general_config、add_image_size_config等函数注册参数,其定义与默认值集中在 diffsynth/diffusion/parsers.py,与官方文档一一对应:
数据集基础配置
| 参数 | 默认值 | 说明 | |-|-|-| |--dataset_base_path| 必填 | 数据集的根目录 | |--dataset_metadata_path|None| 数据集的元数据文件路径 | |--dataset_repeat|1| 每个 epoch 中数据集重复的次数 | |--dataset_num_workers|0| 每个 Dataloader 的进程数量 | |--data_file_keys|"image,video"| 元数据中需要加载的字段名称(通常是图像或视频文件路径),以,分隔 |
模型加载配置
| 参数 | 默认值 | 说明 | |-|-|-| |--model_paths|None| 要加载的模型路径,JSON 格式 | |--model_id_with_origin_paths|None| 带原始文件路径的模型 ID,以,分隔,格式如krea/Krea-2-Raw:raw.safetensors| |--extra_inputs|None| Pipeline 所需的额外输入参数,以,分隔 | |--fp8_models|None| 以 FP8 格式加载的模型,目前仅支持参数不被梯度更新的模型 | |--quant_options|None| 对加载的模型进行动态量化。以;分隔多个条目,每个条目格式为<模型字符串>:<method>[/<exclude_modules>],<模型字符串>需与--model_paths/--model_id_with_origin_paths中的一致,method为已注册的量化方法(如bitsandbytes_nf4),exclude_modules为可选的保持全精度的层 |
训练基础配置
| 参数 | 默认值 | 说明 | |-|-|-| |--learning_rate|1e-4| 学习率 | |--num_epochs|1| 训练轮数(Epoch) | |--trainable_models|None| 可训练的模型,如dit、vae、text_encoder| |--find_unused_parameters|False| DDP 训练中是否查找未使用的参数 | |--weight_decay|0.01| 权重衰减大小 | |--task|"sft"| 训练任务,Krea-2 支持sft、sft:data_process、sft:train|
输出配置
| 参数 | 默认值 | 说明 | |-|-|-| |--output_path|"./models"| 模型保存路径 | |--remove_prefix_in_ckpt|"pipe.dit."| 保存时从 state dict 中移除此前缀 | |--save_steps|None| 保存模型的训练步数间隔;为None时每个 epoch 保存一次 |
LoRA 配置
| 参数 | 默认值 | 说明 | |-|-|-| |--lora_base_model|None| LoRA 添加到哪个模型上 | |--lora_target_modules|"q,k,v,o,ffn.0,ffn.2"| LoRA 添加到哪些层上 | |--lora_rank|32| LoRA 的秩 | |--lora_checkpoint|None| LoRA 检查点路径,提供则从此恢复 LoRA | |--preset_lora_path|None| 预置 LoRA 检查点路径,用于 LoRA 差分训练 | |--preset_lora_model|None| 预置 LoRA 融入的模型,如dit|
梯度配置
| 参数 | 默认值 | 说明 | |-|-|-| |--use_gradient_checkpointing|False| 是否启用梯度检查点,用计算换显存 | |--use_gradient_checkpointing_offload|False| 是否将梯度检查点卸载到内存(CPU) | |--gradient_accumulation_steps|1| 梯度累积步数 |
分辨率配置
| 参数 | 默认值 | 说明 | |-|-|-| |--height|None| 图像高度;留空启用动态分辨率 | |--width|None| 图像宽度;留空启用动态分辨率 | |--max_pixels|1024*1024| 动态分辨率下的最大像素面积,超过此值的图片会被缩小 |
7.3 Krea-2 专有训练参数
除通用参数外,krea2_parser()(train.py)额外注册了三个专有参数:
| 参数 | 说明 | |-|-| |--tokenizer_path| tokenizer 的路径,留空则自动从远程下载(默认 Qwen3-VL-4B-Instruct) | |--initialize_model_on_cpu| 是否在 CPU 上初始化模型,用于进一步降低初始化阶段的显存峰值 | |--align_to_opensource_format| 是否将 LoRA 权重格式对齐为开源格式,便于生成与其他框架兼容的 LoRA 模型;启用后由Krea2LoRAConverter.align_to_opensource_format在保存时转换 state dict(train.py) |
7.4 全量微调示例
以 examples/krea2/model_training/full/Krea-2-Raw.sh 为例,全量微调直接对 DiT 进行 SFT,训练前需先运行accelerate config配置 GPU、DeepSpeed 等环境:
# Please run `accelerate config` to configure GPU, DeepSpeed, etc. modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "krea2/Krea-2-Raw/*" --local_dir ./data/diffsynth_example_dataset accelerate launch examples/krea2/model_training/train.py \ --dataset_base_path data/diffsynth_example_dataset/krea2/Krea-2-Raw \ --dataset_metadata_path data/diffsynth_example_dataset/krea2/Krea-2-Raw/metadata.csv \ --max_pixels 1048576 \ --dataset_repeat 50 \ --model_id_with_origin_paths "krea/Krea-2-Raw:raw.safetensors,Qwen/Qwen3-VL-4B-Instruct:*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors" \ --tokenizer_path "Qwen/Qwen3-VL-4B-Instruct:" \ --learning_rate 1e-5 \ --num_epochs 2 \ --remove_prefix_in_ckpt "pipe.dit." \ --output_path "./models/train/Krea-2-Raw_full" \ --trainable_models "dit" \ --use_gradient_checkpointing \ --find_unused_parameters关键点说明:
- 训练数据通过
--model_id_with_origin_paths一次加载 DiT、文本编码器与 VAE,但仅--trainable_models "dit"参与梯度更新; --dataset_repeat 50将数据集重复 50 次以增加训练步数;- 全量微调建议使用较小学习率(Raw 为
1e-5); - 训练使用流匹配 SFT 损失:
Krea2ImageTrainingModule的task_to_loss将sft映射到FlowMatchSFTLoss(train.py); - 训练时
get_pipeline_inputs将cfg_scale固定为 1、rand_device设为设备,并将图像的原始宽高作为动态分辨率输入(train.py),这与全量脚本中不传--height/--width、只传--max_pixels的动态分辨率策略一致。
Krea-2-Turbo 的全量微调脚本 examples/krea2/model_training/full/Krea-2-Turbo.sh 结构完全一致,仅将模型 ID 与权重文件替换为krea/Krea-2-Turbo:turbo.safetensors。
7.5 LoRA 训练示例
以 examples/krea2/model_training/lora/Krea-2-Raw.sh 为例,LoRA 训练仅需追加--lora_*系列参数:
modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "krea2/Krea-2-Raw/*" --local_dir ./data/diffsynth_example_dataset accelerate launch examples/krea2/model_training/train.py \ --dataset_base_path data/diffsynth_example_dataset/krea2/Krea-2-Raw \ --dataset_metadata_path data/diffsynth_example_dataset/krea2/Krea-2-Raw/metadata.csv \ --max_pixels 1048576 \ --dataset_repeat 50 \ --model_id_with_origin_paths "krea/Krea-2-Raw:raw.safetensors,Qwen/Qwen3-VL-4B-Instruct:*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors" \ --tokenizer_path "Qwen/Qwen3-VL-4B-Instruct:" \ --learning_rate 1e-4 \ --num_epochs 5 \ --remove_prefix_in_ckpt "pipe.dit." \ --output_path "./models/train/Krea-2-Raw_lora" \ --lora_base_model "dit" \ --lora_target_modules "wq,wk,wv,gate,wo,gate,up,down,first,tmlp.0,tmlp.2,projector,txtmlp.1,txtmlp.3,last.linear,tproj.1" \ --lora_rank 32 \ --use_gradient_checkpointing \ --find_unused_parameters \ --align_to_opensource_format要点说明:
--lora_target_modules覆盖了 DiT 中的注意力投影(wq,wk,wv,wo)、MLP(gate,up,down)、时间调制(first)、文本融合层(tmlp.*、txtmlp.*)、投影层(projector、tproj.1)与输出层(last.linear);- LoRA 训练学习率可高于全量微调(此处
1e-4); --align_to_opensource_format让保存的 LoRA 兼容开源社区格式。
关于如何编写模型训练脚本,请参考模型训练;更多高阶训练算法(如分片训练、卸载训练、DeepSpeed、Differential LoRA 等)请参考训练框架详解。
7.6 训练后验证
仓库为全量与 LoRA 训练分别提供了验证脚本:
- 全量验证examples/krea2/model_training/validate_full/Krea-2-Raw.py:加载
models/train/Krea-2-Raw_full/epoch-1.safetensors并通过pipe.dit.load_state_dict(load_state_dict(...))注入后推理; - LoRA 验证examples/krea2/model_training/validate_lora/Krea-2-Raw.py:通过
pipe.load_lora(pipe.dit, "models/train/Krea-2-Raw_lora/epoch-4.safetensors")加载 LoRA 权重后推理,并提示在 Raw 上训练的 LoRA 推荐加载到 Turbo 上使用(validate_lora/Krea-2-Raw.py)。
7.7 高阶:分阶段(Split)训练
仓库还提供了分阶段训练方案 examples/krea2/model_training/special/split_training/Krea-2-Raw.sh,将训练拆成两个阶段以进一步降低显存:
- Stage 1(数据预处理):以
--task sft:data_process运行,将确定性预处理结果(含文本编码等)缓存到磁盘,并通过--offload_models krea/Krea-2-Raw:raw.safetensors把 DiT 卸载,仅保留编码所需模型; - Stage 2(缓存训练):以
--task sft:train运行,--dataset_base_path指向 Stage 1 产出的缓存目录,此时将文本编码器与 VAE 卸载(--offload_models 'Qwen/Qwen3-VL-4B-Instruct:*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors'),训练循环直接读取缓存的嵌入,不再重复编码。
8. 许可协议
⚠️ 提示:Krea-2权重(Raw 与 Turbo)遵循 Krea 2 Community License,不同于DiffSynth-Studio 本身的 Apache 2.0 协议。使用前请务必确认你的应用场景符合该社区许可的条款。
【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-Studio
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考