YAML Metadata Warning:empty or missing yaml metadata in repo card

Check out the documentation for more information.

SDv1.5-DMD2 蒸馏模型

  • FID 22.58

模型简介

本模型基于 Stable Diffusion v1.5 (sdv1.5) 原始模型,采用 DMD2 蒸馏算法 完成轻量化蒸馏训练,在保留SD1.5原生图文生成能力的基础上,优化了模型推理效率、降低了硬件资源消耗,适配日常AI绘图、轻量化部署、批量生成等场景使用。

模型全程基于国内镜像环境训练,解决了原生Hugging Face下载卡顿、连接不稳定问题,可直接在ModelScope平台快速部署、推理与二次微调。

模型基本信息

  • 基础底座:runwayml/stable-diffusion-v1-5

  • 蒸馏算法:DMD2

  • 模型归属:HZNing/SDv1.5-DMD2(ModelScope)

  • 任务类型:文生图(Text-to-Image)

  • 训练分辨率:512×512

  • 潜空间分辨率:64×64

训练硬件与环境配置

硬件设备

训练采用 4张 RTX PRO 6000 (96G) 显卡,分布式并行训练,显存充足,适配大批量数据蒸馏,有效保障模型蒸馏精度与训练稳定性。

环境配置

  • 模型下载镜像:HF-Mirror(https://hf-mirror.com),规避境外网络限制,提升训练资源加载速度

  • 训练框架:PyTorch + torchrun 分布式训练

  • 精度模式:FP16 混合精度训练,兼顾训练速度与模型效果

  • 梯度策略:开启梯度检查点(gradient_checkpointing),降低训练显存占用

核心训练超参数

本次蒸馏训练累计迭代 15000步,预设最大训练迭代200000步,核心参数配置如下:

  • 生成器学习率(generator_lr):1e-5

  • 引导学习率(guidance_lr):1e-5

  • 批次大小(batch_size):32

  • 网格尺寸(grid_size):2

  • 真实引导系数(real_guidance_scale):1.75

  • 虚拟引导系数(fake_guidance_scale):1.0

  • 梯度最大范数(max_grad_norm):10.0

  • 虚假生成更新比例(dfake_gen_update_ratio):10

  • 随机种子(seed):10,保证训练可复现

  • 日志打印间隔:5000步

  • WandB日志记录间隔:50步

训练数据配置

  • 真实图像潜变量数据:sd_vae_latents_laion_500k_lmdb

  • 训练提示词数据集:captions_laion_score6.25.pkl

  • 数据筛选标准:LAION 6.25高分图文对,保证训练数据质量,提升模型生成画面质感与图文匹配度

模型文件说明

本仓库包含完整蒸馏训练产出文件,覆盖模型权重、优化器参数、调度器、训练随机状态等全量文件,支持继续微调、断点续训、完整复现训练过程:

  • .gitattributes:仓库文件属性配置文件,规范文件编码与上传格式

  • model.safetensors:核心蒸馏模型权重文件(安全格式,无恶意代码风险)

  • model_1.safetensors:迭代备份模型权重文件,留存训练中间最优权重

  • optimizer.bin:主优化器参数文件,支持断点续训

  • optimizer_1.bin:备份优化器参数文件

  • random_states_0.pkl:训练随机状态文件,固定随机参数,保障实验完全可复现

  • scheduler.bin:主调度器参数文件,存储学习率调度、迭代策略配置

  • scheduler_1.bin:备份调度器参数文件

ModelScope 使用说明

1. 模型加载推理

可直接通过ModelScope官方接口加载本蒸馏模型,快速实现文生图推理:

import torch
import os
import requests
from safetensors.torch import load_file
from diffusers import UNet2DConditionModel, DiffusionPipeline
from diffusers import DPMSolverMultistepScheduler
import matplotlib.pyplot as plt
from tqdm import tqdm # 用于显示下载进度

# --------------------------
# 1. 配置参数
# --------------------------
base_model_id = "runwayml/stable-diffusion-v1-5"
model_url = "https://modelscope.cn/models/HZNing/SDv1.5-DMD2/resolve/master/model.safetensors"
local_cache_path = "/kaggle/working/temp_dmd2_unet.safetensors" # 临时保存路径

# --------------------------
# 2. 下载模型权重 (如果不存在则下载)
# --------------------------
if not os.path.exists(local_cache_path):
    print(f"📥 开始下载 DMD2 UNet 权重...")
    try:
        response = requests.get(model_url, stream=True, timeout=60)
        response.raise_for_status()
        
        total_size = int(response.headers.get('content-length', 0))
        block_size = 1024 * 1024 # 1MB
        
        with open(local_cache_path, 'wb') as f, tqdm(
            total=total_size, unit='iB', unit_scale=True, desc="Downloading"
        ) as progress_bar:
            for data in response.iter_content(block_size):
                size = f.write(data)
                progress_bar.update(size)
        print("✅ 下载完成!")
    except Exception as e:
        raise RuntimeError(f"下载失败: {e}")
else:
    print("✅ 发现缓存文件,跳过下载。")

# --------------------------
# 3. 加载 UNet 结构 + 本地 safetensors 权重
# --------------------------
print("⏳ 正在加载 UNet 结构...")
# 注意:from_config 只加载结构,不加载权重,速度很快
unet = UNet2DConditionModel.from_config(base_model_id, subfolder="unet")

print("⏳ 正在加载 DMD2 权重...")
unet_weights = load_file(local_cache_path)

# 将权重加载到 unet 中
# strict=False 允许部分键名不匹配(虽然 DMD2 通常完全匹配)
missing_keys, unexpected_keys = unet.load_state_dict(unet_weights, strict=False)
if missing_keys:
    print(f"⚠️ 缺失的键: {missing_keys}")
if unexpected_keys:
    print(f"⚠️ 多余的键: {unexpected_keys}")

# 移动到 GPU 并转换为 FP16
unet = unet.to("cuda", dtype=torch.float16)

# --------------------------
# 4. 组装完整管道
# --------------------------
print("⏳ 正在构建 Diffusion Pipeline...")
pipe = DiffusionPipeline.from_pretrained(
    base_model_id,
    unet=unet,
    torch_dtype=torch.float16,
    variant="fp16",
    safety_checker=None  # 关闭安全检测器提速
).to("cuda")

# 【关键】强制切换为 DPM++ 采样器 (DMD2 推荐少步数推理)
pipe.scheduler = DPMSolverMultistepScheduler.from_config(pipe.scheduler.config)

# --------------------------
# 5. 生成测试
# --------------------------
prompt = "a cute cat sitting on a windowsill, cinematic lighting, 8k"
negative_prompt = "ugly, blurry, low quality, deformed, extra limbs"

inference_steps = 5 # DMD2 的核心优势就是极少步数
generator = torch.Generator(device="cuda").manual_seed(10)

print(f"🚀 开始生成 (Steps: {inference_steps})...")
image = pipe(
    prompt=prompt,
    negative_prompt=negative_prompt,
    num_inference_steps=inference_steps,
    guidance_scale=1.75, # DMD2 通常使用较低的 guidance_scale
    generator=generator
).images[0]

# 展示图片
plt.figure(figsize=(8, 8))
plt.imshow(image)
plt.axis("off")
plt.title(f"DMD2 Result ({inference_steps} steps)")
plt.show()

# 可选:清理临时文件以释放 Kaggle 磁盘空间
# os.remove(local_cache_path)

消融实验

进行了小批量(batch_size=14)图片和大批量(batch_size=432)图片训练对比。以下是不同批量的loss_fake_mean值对比。 小批量(batch_size=1*4)图片和大批量(batch_size=4*32)图片训练对比 从这次实验也可以看出,在分布匹配蒸馏时,模型更易于在大批量图片中学到分布规律,学习更稳定。

开源说明

本模型开源至ModelScope平台(HZNing/SDv1.5-DMD2),仅供学术研究、技术交流与二次开发使用,禁止用于违规商业用途。欢迎各位开发者下载使用、微调优化与技术交流。

Downloads last month

-

Downloads are not tracked for this model. How to track
Safetensors
Model size
0.9B params
Tensor type
F32
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support