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值对比。
从这次实验也可以看出,在分布匹配蒸馏时,模型更易于在大批量图片中学到分布规律,学习更稳定。
开源说明
本模型开源至ModelScope平台(HZNing/SDv1.5-DMD2),仅供学术研究、技术交流与二次开发使用,禁止用于违规商业用途。欢迎各位开发者下载使用、微调优化与技术交流。