mistral.rs 扩散模型图像生成实战:用 Rust 驱动 FLUX.1-schnell 完成文生图推理
【免费下载链接】mistral.rsFast, flexible LLM inference项目地址: https://gitcode.com/GitHub_Trending/mi/mistral.rs
本文基于 mistral.rs 官方示例mistralrs/examples/models/diffusion/main.rs(对应文档 docs/src/content/docs/examples/rust/models/diffusion.md),讲解如何用 Rust 调用DiffusionModelBuilder加载 FLUX 扩散模型并生成图像。读者将掌握扩散模型加载器类型、DiffusionGenerationParams参数、ImageGenerationResponseFormat输出格式等核心 API,并能直接运行一个可复现的文生图程序。
示例概览与运行方式
该示例的目标非常直接:加载black-forest-labs/FLUX.1-schnell模型,输入一段文本提示词,输出一张 720x1280 的图像。运行命令为:
cargo run --release --example diffusion -p mistralrs从源码结构看,示例位于mistralrs/examples/models/diffusion/main.rs,与文档中展示的代码完全一致。整段代码仅约 40 行,覆盖了扩散模型推理的三个核心环节:构建模型(builder)→ 发起生成请求 → 打印结果。
代码开头引入了三个关键类型:
use mistralrs::{ DiffusionGenerationParams, DiffusionLoaderType, DiffusionModelBuilder, ImageGenerationResponseFormat, };DiffusionModelBuilder:用于配置并加载扩散模型;DiffusionLoaderType:指定扩散模型架构的加载方式;DiffusionGenerationParams:控制生成图像的分辨率;ImageGenerationResponseFormat:控制生成结果的返回格式。
模型加载:DiffusionModelBuilder 与 DiffusionLoaderType
构建与加载流程
示例通过以下代码完成模型加载:
let model = DiffusionModelBuilder::new( "black-forest-labs/FLUX.1-schnell", DiffusionLoaderType::FluxOffloaded, ) .with_logging() .build() .await?;DiffusionModelBuilder::new接收两个必填参数:
| 参数 | 说明 |
|---|---|
model_id | Hugging Face 上的模型仓库标识,例如black-forest-labs/FLUX.1-schnell |
loader_type | 扩散模型加载架构,当前支持DiffusionLoaderType::Flux与DiffusionLoaderType::FluxOffloaded |
DiffusionModelBuilder的实现位于 mistralrs/src/diffusion_model.rs。其构造时会应用一组默认值:dtype 为ModelDType::Auto(自动选择)、token 来源为TokenSource::CacheToken(读取~/.cache/huggingface/token)、最大并发序列数为 32、不强制 CPU、关闭日志。
build()内部会调用build_diffusion_pipeline完成 pipeline 的组装,再通过build_model_from_pipeline返回一个可直接使用的Model实例。
可选的配置方法
DiffusionModelBuilder还提供以下链式配置方法(见 diffusion_model.rs):
| 方法 | 作用 |
|---|---|
with_dtype(dtype) | 以指定精度(如ModelDType::F32、ModelDType::BF16)加载模型,默认Auto |
with_force_cpu() | 强制使用 CPU 设备;注释明确指出不要与 PagedAttention 同时使用 |
with_token_source(source) | 指定 Hugging Face token 来源(如TokenSource::CacheToken、TokenSource::Env、TokenSource::Literal) |
with_hf_revision(rev) | 指定 Hugging Face 远端模型的 revision |
with_max_num_seqs(n) | 设置最多同时运行的序列数,默认 32 |
with_logging() | 开启日志输出 |
Flux 与 FluxOffloaded 的区别
DiffusionLoaderType定义于 mistralrs-core/src/pipeline/loaders/diffusion_loaders.rs,目前包含两个变体:
Flux:常规加载,Transformer 与 VAE 全部放入 GPU;FluxOffloaded:将 FLUX Transformer 部分 offload 到 CPU,VAE 仍保留在 GPU 上。
其底层机制可以从 loader 的force_cpu_vb()方法看到:vec![self.offload, false]—— 第一个元素对应 FLUX 权重(offload 时强制 CPU VarBuilder),第二个元素对应 VAE 权重(始终为 false,即留在 GPU)。这在显存不足的环境中非常实用。
自动检测逻辑
DiffusionLoaderType::auto_detect_from_files会根据仓库文件列表自动识别是否为 FLUX 模型:要求存在transformer/config.json、vae/config.json、ae.safetensors,且存在匹配^flux\d+-(schnell|dev)\.safetensors$正则的权重文件。这一逻辑意味着后续新增其他扩散架构(如 SDXL)时,只需扩展该检测函数即可。
在load()实现中,loader 会分别下载 FLUX 权重文件(匹配flux\d+-(schnell|dev).safetensors)与ae.safetensors(自编码器,即 VAE),并校验 FLUX 与 VAE 的 dtype 一致后,构造FluxStepper完成采样。
生成请求:generate_image 与生成参数
请求调用
let response = model .generate_image( "A vibrant sunset in the mountains, 4k, high quality.".to_string(), ImageGenerationResponseFormat::Url, DiffusionGenerationParams::default(), None, ) .await?;Model::generate_image定义于 mistralrs/src/model.rs,签名如下:
pub async fn generate_image( &self, prompt: impl ToString, response_format: ImageGenerationResponseFormat, generation_params: DiffusionGenerationParams, save_file: Option<PathBuf>, ) -> crate::error::Result<ImageGenerationResponse>四个参数分别是:提示词文本、响应格式、生成参数、可选的文件保存路径。此外还有generate_image_with_model变体,多一个model_id: Option<&str>参数,用于在加载了多个模型时指定使用哪个模型生成(None表示使用默认模型)。
生成参数 DiffusionGenerationParams
DiffusionGenerationParams定义于 mistralrs-core/src/diffusion_models/mod.rs:
pub struct DiffusionGenerationParams { pub height: usize, pub width: usize, }目前只有height和width两个字段,控制输出图像的分辨率。Default实现固定为720x1280(竖版比例)。如果需要生成横向或其他比例的图像,可以手动构造:
use mistralrs::DiffusionGenerationParams; let params = DiffusionGenerationParams { height: 1024, width: 1024, };响应格式 ImageGenerationResponseFormat
该枚举定义于 mistralrs-core/src/request.rs,包含两个变体:
| 变体 | 说明 |
|---|---|
ImageGenerationResponseFormat::Url | 图像保存为文件后返回 URL 字符串 |
ImageGenerationResponseFormat::B64Json | 返回 Base64 编码的图像数据(JSON 格式) |
在Request::Normal的RequestMessage::ImageGeneration变体(request.rs)中,请求会携带prompt、format、generation_params与可选的save_file字段,随后进入 pipeline 由DiffusionModel::forward真正执行采样。
结果处理:耗时统计与输出
生成完成后,示例统计耗时并打印图像保存位置:
let finished = Instant::now(); println!( "Done! Took {} s. Image saved at: {}", finished.duration_since(start).as_secs_f32(), response.data[0].url.as_ref().unwrap() );response.data[0].url是Option<String>,当使用Url格式且生成成功时包含图像的本地文件路径;若使用B64Json格式,则相应字段存放 Base64 数据。使用as_ref().unwrap()之前建议先判空,避免生成失败时 panic。
底层原理:FluxStepper 与采样参数
生成过程的核心是 mistralrs-core/src/diffusion_models/flux/stepper.rs 中的FluxStepper。它会同时加载两个文本编码器(T5 与 CLIP)对提示词编码,再由 FLUX Transformer 与 VAE 自编码器完成去噪与解码。
采样步数由FluxStepperConfig::default_for_guidance(stepper.rs)根据模型是否使用 guidance 自动决定:
- 有 guidance 的模型(如 FLUX.1-dev):默认 50 步,并启用
FluxStepperShift(base_shift: 0.5、max_shift: 1.15、guidance_scale: 4.0); - 无 guidance 的模型(如 FLUX.1-schnell):默认仅4 步,不启用 guidance。
这也是 FLUX.1-schnell 主打快速推理的原因——以更少的采样步数换取更快的生成速度。FluxStepper::new后续会调用flux::sampling::get_schedule根据num_steps生成时间步调度。
Python 等价实现
mistral.rs 同时提供 Python 绑定(mistralrs-pyo3),对应示例为 examples/python/flux.py,接口语义与 Rust 版一一对应:
from mistralrs import ( Runner, Which, DiffusionArchitecture, ImageGenerationResponseFormat, ) runner = Runner( which=Which.DiffusionPlain( model_id="black-forest-labs/FLUX.1-schnell", arch=DiffusionArchitecture.FluxOffloaded, ), ) res = runner.generate_image( "A vibrant sunset in the mountains, 4k, high quality.", ImageGenerationResponseFormat.Url, ) print(res.data[0].url)其中DiffusionArchitecture.FluxOffloaded与 Rust 侧的DiffusionLoaderType::FluxOffloaded对应,generate_image的默认分辨率同样是 720x1280。
注意事项
- 显存:FLUX 系列模型体积较大,显存有限时优先使用
FluxOffloaded加载方式,将 Transformer 部分 offload 至 CPU; - token 认证:下载 gated 模型需要有效的 Hugging Face token,默认从
~/.cache/huggingface/token读取,可通过with_token_source更换来源; - 分辨率:
DiffusionGenerationParams::default()为 720x1280,修改分辨率会影响生成耗时与显存占用; - 保存文件:
generate_image的save_file参数可指定图像落盘路径,方便在批处理场景下管理输出。
相关参考文件:
- 示例源码:mistralrs/examples/models/diffusion/main.rs
- Builder 实现:mistralrs/src/diffusion_model.rs
- 加载器与
DiffusionLoaderType:mistralrs-core/src/pipeline/loaders/diffusion_loaders.rs - 生成参数定义:mistralrs-core/src/diffusion_models/mod.rs
- 采样核心:mistralrs-core/src/diffusion_models/flux/stepper.rs
- Python 等价示例:examples/python/flux.py
【免费下载链接】mistral.rsFast, flexible LLM inference项目地址: https://gitcode.com/GitHub_Trending/mi/mistral.rs
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考