跳到正文
北京时间
原文
Hugging Face:Blog·· 2025-06-19精选AI 评分69

Hugging Face 教程:用 QLoRA 在 RTX 4090 上以约 9GB 显存微调 FLUX.1-dev

(LoRA) Fine-Tuning FLUX.1-dev on Consumer Hardware

AI 导读

Hugging Face 发布教程,讲解用 QLoRA 和 diffusers 库在单卡消费级硬件上微调 FLUX.1-dev,学习 Alphonse Mucha 画风。

推荐理由

原文给出完整训练配置、显存与耗时数据,读者可以照着在消费级 GPU 上复现 FLUX.1-dev 的 QLoRA 微调。

正文 · AI 翻译

Open In Colab

在上一篇文章探索 Diffusers 中的量化后端中,我们深入探讨了各种量化技术如何缩小 FLUX.1-dev 等扩散模型,使其在推理时更易于使用,而不会大幅降低性能。我们看到了 bitsandbytes、torchao 等方法如何减少生成图像时的内存占用。

进行推理很酷,但要让这些模型真正成为我们自己的,我们还需要能够对它们进行微调。因此,在本文中,我们将解决这些模型的高效微调问题,在单块 GPU 上峰值显存使用控制在约 10 GB 以下。本文将指导你使用 diffusers 库通过 QLoRA 微调 FLUX.1-dev。我们将展示在 NVIDIA RTX 4090 上的结果。我们还将重点介绍使用 torchao 进行 FP8 训练如何在兼容硬件上进一步优化速度。

目录

数据集

我们的目标是微调 black-forest-labs/FLUX.1-dev,使其采用 Alphonse Mucha 的艺术风格,使用一个小型数据集。

FLUX 架构

该模型由三个主要组件组成:

  • 文本编码器(CLIP 和 T5)
  • Transformer(主模型 - Flux Transformer)
  • 变分自编码器(VAE)

在我们的 QLoRA 方法中,我们仅专注于微调transformer 组件。文本编码器和 VAE 在整个训练过程中保持冻结。

使用 Diffusers 对 FLUX.1-dev 进行 QLoRA 微调

我们使用了 diffusers 训练脚本(稍作修改自此处,专为 FLUX 模型的 DreamBooth 风格 LoRA 微调而设计。此外,一个用于复现本博文结果(并在 Google Colab 中使用)的简化版本可在此处获取。让我们来看看 QLoRA 和内存效率的关键部分:

关键优化技术

LoRA(低秩适应)深入解析:LoRA 通过使用低秩矩阵跟踪权重更新,使模型训练更加高效。LoRA 不是更新完整的权重矩阵 W W ,而是学习两个较小的矩阵 A A 和 B B 。模型权重的更新为 ΔW=BA \Delta W = B A ,其中 A∈Rr×k A \in \mathbb{R}^{r \times k} 和 B∈Rd×r B \in \mathbb{R}^{d \times r} 。数字 r r (称为秩)远小于原始维度,这意味着需要更新的参数更少。最后,α \alpha 是 LoRA 激活的缩放因子。它影响 LoRA 对更新的影响程度,通常设置为与 r r 相同或为其倍数。它有助于平衡预训练模型和 LoRA 适配器的影响。有关该概念的通用介绍,请查看我们之前的博文:使用 LoRA 进行高效 Stable Diffusion 微调。

Illustration of LoRA injecting two low-rank matrices around a frozen weight matrix

QLoRA:效率利器: QLoRA 通过首先以量化格式(通常通过 bitsandbytes 使用 4 位)加载预训练基础模型来增强 LoRA,大幅削减基础模型的内存占用。然后,它在这个量化基础模型之上训练 LoRA 适配器(通常为 FP16/BF16)。这显著降低了保存基础模型所需的显存。

例如,在 HiDream 的 DreamBooth 训练脚本 中,使用 bitsandbytes 进行 4 位量化可将 LoRA 微调的峰值内存使用量从约 60GB 降至约 37GB,且质量下降可忽略不计甚至没有。正是同样的原理,我们在此将其应用于在消费级硬件上微调 FLUX.1。

8 位优化器(AdamW): 标准 AdamW 优化器以 32 位(FP32)为每个参数维护一阶和二阶矩估计,这会消耗大量内存。8 位 AdamW 使用分块量化以 8 位精度存储优化器状态,同时保持训练稳定性。与标准 FP32 AdamW 相比,该技术可将优化器内存使用量减少约 75%。在脚本中启用它非常简单:


# Check for the --use_8bit_adam flag
if args.use_8bit_adam:
    optimizer_class = bnb.optim.AdamW8bit
else:
    optimizer_class = torch.optim.AdamW

optimizer = optimizer_class(
    params_to_optimize,
    betas=(args.adam_beta1, args.adam_beta2),
    weight_decay=args.adam_weight_decay,
    eps=args.adam_epsilon,
)

梯度检查点: 在前向传播过程中,通常会存储中间激活值用于反向传播的梯度计算。梯度检查点通过仅存储某些检查点激活值并在反向传播期间重新计算其他激活值,以计算换内存。

if args.gradient_checkpointing:
    transformer.enable_gradient_checkpointing()

缓存潜变量: 这种优化技术在训练开始前,通过 VAE 编码器预处理所有训练图像。它将生成的潜变量表示存储在内存中。在训练期间,直接使用缓存的潜变量,而不是实时编码图像。这种方法提供两个主要好处:

  1. 消除了训练期间冗余的 VAE 编码计算,加速每个训练步骤
  2. 允许在缓存后从 GPU 内存中完全移除 VAE。代价是增加 RAM 使用量以存储所有缓存的潜变量,但这对于小型数据集通常是可以管理的。
# Cache latents before training if the flag is set
    if args.cache_latents:
        latents_cache = []
        for batch in tqdm(train_dataloader, desc="Caching latents"):
            with torch.no_grad():
                batch["pixel_values"] = batch["pixel_values"].to(
                    accelerator.device, non_blocking=True, dtype=weight_dtype
                )
                latents_cache.append(vae.encode(batch["pixel_values"]).latent_dist)
        # VAE is no longer needed, free up its memory
        del vae
        free_memory()

设置 4 位量化(BitsAndBytesConfig):

本节演示基础模型的 QLoRA 配置:

# Determine compute dtype based on mixed precision
bnb_4bit_compute_dtype = torch.float32
if args.mixed_precision == "fp16":
    bnb_4bit_compute_dtype = torch.float16
elif args.mixed_precision == "bf16":
    bnb_4bit_compute_dtype = torch.bfloat16

nf4_config = BitsAndBytesConfig(
    load_in_4bit=True,
    bnb_4bit_quant_type="nf4",
    bnb_4bit_compute_dtype=bnb_4bit_compute_dtype,
)

transformer = FluxTransformer2DModel.from_pretrained(
    args.pretrained_model_name_or_path,
    subfolder="transformer",
    quantization_config=nf4_config,
    torch_dtype=bnb_4bit_compute_dtype,
)
# Prepare model for k-bit training
transformer = prepare_model_for_kbit_training(transformer, use_gradient_checkpointing=False)
# Gradient checkpointing is enabled later via transformer.enable_gradient_checkpointing() if arg is set

定义 LoRA 配置(LoraConfig): 将适配器添加到量化后的 transformer:

transformer_lora_config = LoraConfig(
    r=args.rank,
    lora_alpha=args.rank, 
    init_lora_weights="gaussian",
    target_modules=["to_k", "to_q", "to_v", "to_out.0"], # FLUX attention blocks
)
transformer.add_adapter(transformer_lora_config)
print(f"trainable params: {transformer.num_parameters(only_trainable=True)} || all params: {transformer.num_parameters()}")
# trainable params: 4,669,440 || all params: 11,906,077,760

只有这些 LoRA 参数变为可训练。

预计算文本嵌入(CLIP/T5)

在启动 QLoRA 微调之前,我们可以通过一次性缓存文本编码器的输出来节省大量显存和实际时间。

在训练时,数据加载器只需读取缓存的嵌入,而不是重新编码标题,因此 CLIP/T5 编码器永远不必驻留在 GPU 内存中。

代码

# https://github.com/huggingface/diffusers/blob/main/examples/research_projects/flux_lora_quantization/compute_embeddings.py
import argparse

import pandas as pd
import torch
from datasets import load_dataset
from huggingface_hub.utils import insecure_hashlib
from tqdm.auto import tqdm
from transformers import T5EncoderModel

from diffusers import FluxPipeline


MAX_SEQ_LENGTH = 77
OUTPUT_PATH = "embeddings.parquet"


def generate_image_hash(image):
    return insecure_hashlib.sha256(image.tobytes()).hexdigest()


def load_flux_dev_pipeline():
    id = "black-forest-labs/FLUX.1-dev"
    text_encoder = T5EncoderModel.from_pretrained(id, subfolder="text_encoder_2", load_in_8bit=True, device_map="auto")
    pipeline = FluxPipeline.from_pretrained(
        id, text_encoder_2=text_encoder, transformer=None, vae=None, device_map="balanced"
    )
    return pipeline


@torch.no_grad()
def compute_embeddings(pipeline, prompts, max_sequence_length):
    all_prompt_embeds = []
    all_pooled_prompt_embeds = []
    all_text_ids = []
    for prompt in tqdm(prompts, desc="Encoding prompts."):
        (
            prompt_embeds,
            pooled_prompt_embeds,
            text_ids,
        ) = pipeline.encode_prompt(prompt=prompt, prompt_2=None, max_sequence_length=max_sequence_length)
        all_prompt_embeds.append(prompt_embeds)
        all_pooled_prompt_embeds.append(pooled_prompt_embeds)
        all_text_ids.append(text_ids)

    max_memory = torch.cuda.max_memory_allocated() / 1024 / 1024 / 1024
    print(f"Max memory allocated: {max_memory:.3f} GB")
    return all_prompt_embeds, all_pooled_prompt_embeds, all_text_ids


def run(args):
    dataset = load_dataset("Norod78/Yarn-art-style", split="train")
    image_prompts = {generate_image_hash(sample["image"]): sample["text"] for sample in dataset}
    all_prompts = list(image_prompts.values())
    print(f"{len(all_prompts)=}")

    pipeline = load_flux_dev_pipeline()
    all_prompt_embeds, all_pooled_prompt_embeds, all_text_ids = compute_embeddings(
        pipeline, all_prompts, args.max_sequence_length
    )

    data = []
    for i, (image_hash, _) in enumerate(image_prompts.items()):
        data.append((image_hash, all_prompt_embeds[i], all_pooled_prompt_embeds[i], all_text_ids[i]))
    print(f"{len(data)=}")

    # Create a DataFrame
    embedding_cols = ["prompt_embeds", "pooled_prompt_embeds", "text_ids"]
    df = pd.DataFrame(data, columns=["image_hash"] + embedding_cols)
    print(f"{len(df)=}")

    # Convert embedding lists to arrays (for proper storage in parquet)
    for col in embedding_cols:
        df[col] = df[col].apply(lambda x: x.cpu().numpy().flatten().tolist())

    # Save the dataframe to a parquet file
    df.to_parquet(args.output_path)
    print(f"Data successfully serialized to {args.output_path}")


if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument(
        "--max_sequence_length",
        type=int,
        default=MAX_SEQ_LENGTH,
        help="Maximum sequence length to use for computing the embeddings. The more the higher computational costs.",
    )
    parser.add_argument("--output_path", type=str, default=OUTPUT_PATH, help="Path to serialize the parquet file.")
    args = parser.parse_args()

    run(args)

如何使用

python compute_embeddings.py \
  --max_sequence_length 77 \
  --output_path embeddings_alphonse_mucha.parquet

通过将此与缓存的 VAE 潜变量(--cache_latents)结合,您可以将活动模型缩减为仅量化后的 transformer + LoRA 适配器,使整个微调过程轻松保持在 10 GB 显存以下。

设置与结果

在此演示中,我们利用 NVIDIA RTX 4090(24GB 显存)来探索其性能。使用 accelerate 的完整训练命令如下所示。

# You need to pre-compute the text embeddings first. See the diffusers repo.
# https://github.com/huggingface/diffusers/tree/main/examples/research_projects/flux_lora_quantization
accelerate launch --config_file=accelerate.yaml \
  train_dreambooth_lora_flux_miniature.py \
  --pretrained_model_name_or_path="black-forest-labs/FLUX.1-dev" \
  --data_df_path="embeddings_alphonse_mucha.parquet" \
  --output_dir="alphonse_mucha_lora_flux_nf4" \
  --mixed_precision="bf16" \
  --use_8bit_adam \
  --weighting_scheme="none" \
  --width=512 \
  --height=768 \
  --train_batch_size=1 \
  --repeats=1 \
  --learning_rate=1e-4 \
  --guidance_scale=1 \
  --report_to="wandb" \
  --gradient_accumulation_steps=4 \
  --gradient_checkpointing \ # can drop checkpointing when HW has more than 16 GB.
  --lr_scheduler="constant" \
  --lr_warmup_steps=0 \
  --cache_latents \
  --rank=4 \
  --max_train_steps=700 \
  --seed="0"

RTX 4090 的配置: 在我们的 RTX 4090 上,我们使用了 train_batch_size 为 1,gradient_accumulation_steps 为 4,mixed_precision="bf16",gradient_checkpointing=True,use_8bit_adam=True,LoRA rank 为 4,分辨率为 512x768。潜变量使用 cache_latents=True 缓存。

内存占用(RTX 4090):

  • QLoRA: QLoRA 微调的峰值显存使用量约为 9GB。
  • BF16 LoRA:在相同设置下运行标准 LoRA(基础 FLUX.1-dev 使用 FP16),消耗了 26 GB 显存。
  • BF16 全量微调:在没有内存优化的情况下,估计需要约 120 GB 显存。

训练时间(RTX 4090): 在 RTX 4090 上,使用 train_batch_size 为 1、分辨率为 512x768,对 Alphonse Mucha 数据集进行 700 步微调大约耗时 41 分钟。

输出质量: 最终的衡量标准是生成的艺术作品。以下是我们基于 derekl35/alphonse-mucha-style 数据集微调的 QLoRA 模型生成的样本:

此表比较了主要的 bf16 精度结果。微调的目标是让模型学习 Alphonse Mucha 的独特风格。

提示词 基础模型输出 QLoRA 微调输出(Mucha 风格)
“宁静的黑发女子,月光下的百合,盘旋的植物纹样,alphonse mucha style” Base model output for the first prompt QLoRA model output for the first prompt
“池塘里的小狗,alphonse mucha style” Base model output for the second prompt QLoRA model output for the second prompt
“装饰华丽的狐狸,戴着秋叶和浆果的项圈,置身于森林枝叶的织锦中,alphonse mucha style” Base model output for the third prompt QLoRA model output for the third prompt

微调后的模型很好地捕捉了 Mucha 标志性的新艺术风格,这在装饰图案和独特的调色板中显而易见。QLoRA 过程在学习新风格的同时保持了出色的保真度。

点击查看 fp16 对比

结果几乎完全相同,表明 QLoRA 在 fp16 和 bf16 混合精度下都能有效运行。

模型对比:基础模型 vs. QLoRA 微调(fp16)

提示词 基础模型输出 QLoRA 微调输出(Mucha 风格)
“宁静的黑发女子,月光下的百合,盘旋的植物纹样,alphonse mucha style” Base model output for the first prompt QLoRA model output for the first prompt
“池塘里的小狗,alphonse mucha style” Base model output for the second prompt QLoRA model output for the second prompt
“装饰华丽的狐狸,戴着秋叶和浆果的项圈,置身于森林枝叶的织锦中,alphonse mucha style” Base model output for the third prompt QLoRA model output for the third prompt

使用 TorchAO 进行 FP8 微调

对于计算能力为 8.9 或更高的 NVIDIA GPU(例如 H100、RTX 4090)用户,可以通过 torchao 库利用 FP8 训练来实现更高的速度效率。

我们在 H100 SXM GPU 上微调了 FLUX.1-dev LoRA,略微修改了diffusers-torchao 训练脚本。使用了以下命令:

accelerate launch train_dreambooth_lora_flux.py \
  --pretrained_model_name_or_path=black-forest-labs/FLUX.1-dev \
  --dataset_name=derekl35/alphonse-mucha-style --instance_prompt="a woman, alphonse mucha style" --caption_column="text" \
  --output_dir=alphonse_mucha_fp8_lora_flux \
  --mixed_precision=bf16 --use_8bit_adam \
  --weighting_scheme=none \
  --height=768 --width=512 --train_batch_size=1 --repeats=1 \
  --learning_rate=1e-4 --guidance_scale=1 --report_to=wandb \
  --gradient_accumulation_steps=1 --gradient_checkpointing \
  --lr_scheduler=constant --lr_warmup_steps=0 --rank=4 \
  --max_train_steps=700 --checkpointing_steps=600 --seed=0 \
  --do_fp8_training --push_to_hub

训练运行的峰值内存使用量为 36.57 GB,并在大约 20 分钟内完成。

此 FP8 微调模型的定性结果也可查看: FP8 model outputs

使用 torchao 启用 FP8 训练的关键步骤包括:

  1. 注入 FP8 层到模型中,使用来自 torchao.float8 的 convert_to_float8_training。
  2. 定义 module_filter_fn以指定哪些模块应转换为 FP8,哪些不应转换。

如需更详细的指南和代码片段,请参阅此 gist 和 diffusers-torchao 仓库。

使用训练好的 LoRA 适配器进行推理

训练好你的 LoRA 适配器后,有两种主要的推理方法。

选项 1:加载 LoRA 适配器

一种方法是在基础模型之上加载训练好的 LoRA 适配器。

加载 LoRA 的好处:

  • 灵活性:无需重新加载基础模型即可轻松切换不同的 LoRA 适配器
  • 实验性:通过更换适配器测试多种艺术风格或概念
  • 模块化:使用 set_adapters() 组合多个 LoRA 适配器以实现创意混合
  • 存储效率:保留单个基础模型和多个小型适配器文件

代码

from diffusers import FluxPipeline, FluxTransformer2DModel, BitsAndBytesConfig
import torch 

ckpt_id = "black-forest-labs/FLUX.1-dev"
pipeline = FluxPipeline.from_pretrained(
    ckpt_id, torch_dtype=torch.float16
)
pipeline.load_lora_weights("derekl35/alphonse_mucha_qlora_flux", weight_name="pytorch_lora_weights.safetensors")

pipeline.enable_model_cpu_offload()

image = pipeline(
    "a puppy in a pond, alphonse mucha style", num_inference_steps=28, guidance_scale=3.5, height=768, width=512, generator=torch.manual_seed(0)
).images[0]
image.save("alphonse_mucha.png")

选项2:将LoRA合并到基础模型

当你想要以单一风格实现最大效率时,你可以合并LoRA权重到基础模型中。

合并LoRA的好处:

  • 显存效率:推理期间没有来自适配器权重的额外内存开销
  • 速度:推理速度稍快,因为无需应用适配器计算
  • 量化兼容性:可以重新量化合并后的模型以实现最大内存效率

代码

from diffusers import FluxPipeline, AutoPipelineForText2Image, FluxTransformer2DModel, BitsAndBytesConfig
import torch 

ckpt_id = "black-forest-labs/FLUX.1-dev"
pipeline = FluxPipeline.from_pretrained(
    ckpt_id, text_encoder=None, text_encoder_2=None, torch_dtype=torch.float16
)
pipeline.load_lora_weights("derekl35/alphonse_mucha_qlora_flux", weight_name="pytorch_lora_weights.safetensors")
pipeline.fuse_lora()
pipeline.unload_lora_weights()

pipeline.transformer.save_pretrained("fused_transformer")

bnb_4bit_compute_dtype = torch.bfloat16

nf4_config = BitsAndBytesConfig(
    load_in_4bit=True,
    bnb_4bit_quant_type="nf4",
    bnb_4bit_compute_dtype=bnb_4bit_compute_dtype,
)
transformer = FluxTransformer2DModel.from_pretrained(
    "fused_transformer",
    quantization_config=nf4_config,
    torch_dtype=bnb_4bit_compute_dtype,
)

pipeline = AutoPipelineForText2Image.from_pretrained(
    ckpt_id, transformer=transformer, torch_dtype=bnb_4bit_compute_dtype
)
pipeline.enable_model_cpu_offload()

image = pipeline(
    "a puppy in a pond, alphonse mucha style", num_inference_steps=28, guidance_scale=3.5, height=768, width=512, generator=torch.manual_seed(0)
).images[0]
image.save("alphonse_mucha_merged.png")

在Google Colab上运行

虽然我们在RTX 4090上展示了结果,但相同的代码可以在更易获取的硬件上运行,例如Google Colab中免费提供的T4 GPU。

在T4上,你可以预期微调过程会显著更长,相同步数大约需要4小时。这是为了可访问性所做的权衡,但它使得无需高端硬件即可进行自定义微调。如果在Colab上运行,请注意使用限制,因为4小时的训练运行可能会触及这些限制。

结论

QLoRA与diffusers库相结合,显著地民主化了定制最先进模型(如FLUX.1-dev)的能力。正如在RTX 4090上展示的那样,高效微调触手可及,能够产生高质量的样式适配。此外,对于拥有最新NVIDIA硬件的用户,torchao通过FP8精度实现了更快的训练。

在Hub上分享你的创作!

分享你微调后的LoRA适配器是为开源社区做贡献的绝佳方式。它允许他人轻松尝试你的风格,基于你的工作进行构建,并有助于创建一个充满活力的创意AI工具生态系统。

如果你为FLUX.1-dev训练了一个LoRA,我们鼓励你分享它。最简单的方法是在训练脚本中添加--push_to_hub标志。或者,如果你已经训练了一个模型并想上传它,可以使用以下代码片段。

# Prereqs:
# - pip install huggingface_hub diffusers
# - Run `huggingface-cli login` (or set HF_TOKEN env-var) once.
# - save model

from huggingface_hub import create_repo, upload_folder

repo_id = "your-username/alphonse_mucha_qlora_flux"
create_repo(repo_id, exist_ok=True)

upload_folder(
    repo_id=repo_id,
    folder_path="alphonse_mucha_qlora_flux",
    commit_message="Add Alphonse Mucha LoRA adapter"
)

查看我们的Mucha LoRA和TorchAO FP8 LoRA。你可以在这个集合中找到两者以及其他适配器。

我们迫不及待地想看到你创造的作品!

来源:Hugging Face:Blog · huggingface.co