资讯详情

Krea-2 图像生成模型实战指南:在 DiffSynth-Studio 中完成推理、低显存部署与全量/LoRA 训练

📅 2026/9/15 18:27:33 | 华诺云谱 👁 阅读
Krea-2 图像生成模型实战指南:在 DiffSynth-Studio 中完成推理、低显存部署与全量/LoRA 训练
Krea-2 图像生成模型实战指南在 DiffSynth-Studio 中完成推理、低显存部署与全量/LoRA 训练【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-StudioKrea-2 是 Krea 团队开发的图像生成模型本指南以 DiffSynth-Studio 仓库中的 Krea-2 官方文档 为主体完整讲解 Krea-2-Raw 与 Krea-2-Turbo 两个版本在 DiffSynth-Studio 中的安装、快速推理、低显存部署、全量微调与 LoRA 训练全流程。读完本文你将掌握Krea2Pipeline的加载与调用方式、全部推理与训练参数的含义与默认值、显存管理配置的底层原理并能够直接复用仓库中提供的示例脚本完成从数据集准备到模型验证的完整工作流。1. Krea-2 与 DiffSynth-StudioKrea-2 是 Krea 团队发布的图像生成模型DiffSynth-Studio 为其提供了完整的推理与训练支持。从源码结构看Krea-2 的推理由 diffsynth/pipelines/krea2.py 中的Krea2Pipeline承载模型整体由三部分组成文本编码器Qwen3-VL-4B-Instruct多模态大语言模型负责将提示词编码为多层的 hidden states对应源码 diffsynth/models/krea2_text_encoder.py去噪网络 DiTSingleStreamDiT单流混合模态 DiT对应源码 diffsynth/models/krea2_dit.py图像 VAEQwen-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-Studiogit 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_dtypetorch.bfloat16, devicecuda, model_configs[ ModelConfig(model_idkrea/Krea-2-Raw, origin_file_patternraw.safetensors, **vram_config), ModelConfig(model_idQwen/Qwen3-VL-4B-Instruct, origin_file_pattern*.safetensors, **vram_config), ModelConfig(model_idQwen/Qwen-Image, origin_file_patternvae/diffusion_pytorch_model.safetensors, **vram_config), ], tokenizer_configModelConfig(model_idQwen/Qwen3-VL-4B-Instruct, origin_file_pattern), vram_limittorch.cuda.mem_get_info(cuda)[1] / (1024 ** 3) - 1, ) prompt A cat standing on a stone. image pipe(prompt, seed0, num_inference_steps52, cfg_scale4.5) image.save(image.jpg)结合Krea2Pipeline.from_pretrained的实现krea2.py这段代码的加载过程可以做如下拆解三个ModelConfig分别对应 Krea-2 模型的三个组件raw.safetensorsDiT 权重、Qwen3-VL-4B-Instruct 的全部 safetensors文本编码器、Qwen-Image 的 VAE 权重。origin_file_pattern用于在远端模型仓库中精确匹配需要下载的权重文件避免下载无关文件tokenizer_config单独指定了 Qwen3-VL-4B-Instruct 的 tokenizerfrom_pretrained内部会先download_if_necessary()再通过AutoTokenizer.from_pretrained(..., max_length512)实例化加载完成后text_encoder、dit、vae会从模型池中按名称取出fetch_model(krea2_text_encoder)等pipe.vram_management_enabled会根据配置自动判断显存管理是否生效。vram_config中的六个键对应显存管理的六个生命周期阶段其中offload_*控制参数卸载后的存储格式与设备此处为磁盘上的disk类型最大限度释放显存onload_*控制参数加载回内存时的格式FP8 可减半内存占用preparing_*控制参数送入 GPU 前的准备阶段FP8computation_*则指定实际参与前向计算的精度与设备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_idQwen/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_steps8、cfg_scale1、mu1.15是固定参数不应随意改动from diffsynth.pipelines.krea2 import Krea2Pipeline, ModelConfig import torch pipe Krea2Pipeline.from_pretrained( torch_dtypetorch.bfloat16, devicecuda, model_configs[ ModelConfig(model_idkrea/Krea-2-Turbo, origin_file_patternturbo.safetensors), ModelConfig(model_idQwen/Qwen3-VL-4B-Instruct, origin_file_pattern*.safetensors), ModelConfig(model_idQwen/Qwen-Image, origin_file_patternvae/diffusion_pytorch_model.safetensors), ], tokenizer_configModelConfig(model_idQwen/Qwen3-VL-4B-Instruct, origin_file_pattern), ) prompt Portrait of a woman in a blue dress, underwater, surrounded by colorful bubbles. image pipe( prompt, seed0, height2048, width2048, # The following parameters are fixed. num_inference_steps8, cfg_scale1, mu1.15, ) image.save(image.jpg)注意Turbo 的cfg_scale1意味着不使用 CFG 引导正向与负向分支权重相同mu1.15则通过流匹配调度器的动态位移机制调整时间步分布。该参数会在self.scheduler.set_timesteps(num_inference_steps, denoising_strength1.0, dynamic_shift_len(height // 16) * (width // 16), mumu)中被消费krea2.py其中dynamic_shift_len与图像分辨率成正比即分辨率越高时间步位移越明显。5.4 推理流程的内部机制源码级从源码结构看Krea2Pipeline.__call__的执行流程可以划分为三个阶段阶段一预处理单元链。推理依次执行五个PipelineUnitkrea2.pyKrea2Unit_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为 512krea2.pyKrea2Unit_InputImageEmbedder当传入input_image图生图场景时用 VAE 编码输入图像并将初始噪声按首个时间步叠加到输入潜变量上否则直接用纯噪声初始化Krea2Unit_PromptEmbPreCompute将文本嵌入通过 DiT 的txtfusion与txtmlp预计算为融合后的条件避免在每个去噪步重复计算——这由context_pre_computeTrue控制。阶段二循环去噪。在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_nf4exclude_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 dicttrain.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映射到FlowMatchSFTLosstrain.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、MLPgate,up,down、时间调制first、文本融合层tmlp.*、txtmlp.*、投影层projector、tproj.1与输出层last.linearLoRA 训练学习率可高于全量微调此处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),仅供参考
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。