DiffSplat
将预训练2D扩散模型改造为3D Gaussian Splat生成器,1~2秒文/图生3D,ICLR
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
将预训练2D扩散模型改造为3D Gaussian Splat生成器,1~2秒文/图生3D,ICLR
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
GitHub: chenguolin/DiffSplat | ⭐ 525 Stars | ICLR 2025
你有没有过这样的经历:看到一张玩具照片,想立刻生成一个可以在 Three.js 里旋转把玩的 3D 模型?或者输入一句"一只骑滑板的兔子",想要一个完整的、可360度浏览的3D资产?
传统的 3D 生成方法面临两座大山:高质量 3D 数据稀缺(互联网上 2D 图片多如牛毛,3D 模型却凤毛麟角)和 多视角一致性差(先生成 4 张视角图再拼成 3D,不同视角之间经常"撞脸")。主流的 Reconstruction 类方法虽然借助 2D 先验做多视角生成再重建 3D,但多步骤流程引入的累计误差让质量不稳定。
DiffSplat 的核心洞察是:与其把扩散模型用来生成多视角图片,不如直接让它生成 3D Gaussian 的参数。Gaussian Splatting(3DGS)是一种用数百万个椭球(Gaussian)来表示场景的技术,每个椭球携带位置、颜色、透明度等信息,可以实时渲染出逼真的新视角图像。DiffSplat 的做法是:把预训练的文生图扩散模型的 U-Net 解码器替换为 Gaussian Splat 解码器,冻结权重并只训练新增的 Splatting 头,让扩散模型强大的 2D 先验知识直接为 3D 服务。
效果惊艳:从一句文字描述或一张图片出发,在 1~2 秒内直接生成完整的 3D Gaussian Splats,无须多步优化,无须多视角输入。在 A100 GPU 上,推理仅需约 2GB VRAM(SD1.5 版本),生成结果可直接导入 WebGL/Three.js 进行实时渲染。
DiffSplat 的架构分为三个层次:
① 结构化 Splat 表示(Structured Splat Representation)
输入是一张 RGB 图片 + 对应的相机位姿(Plucker 坐标),通过 Image Tokenizer(基于 DiagonalGaussianDistribution)将图片编码为 latent 空间向量,再通过基于 Transformer 的 GSRecon 模型将这些隐向量解码为每像素对应的 Gaussian 参数:颜色(3通道)、位置偏移(3通道)、缩放(3通道)、旋转四元数(4通道)、不透明度(1通道)。每个 patch(默认 8×8 像素)对应一个 Gaussian Primitive,全部预测完后通过 Splatting Rasterizer 渲染出新视角的 RGB 图片和深度图。
GSRecon 采用了 LLaMA 风格的 Transformer 架构(12 层,512 维,8 头),使用 patchify/unpatchify 操作实现图像到 token 序列的转换,支持配置渐变检查点(gradient checkpointing)以节省显存。渲染部分基于 3D Gaussian Splatting 的高效光栅化技术,支持 deferred 分块渲染(deferred_bp)以加速大分辨率场景。
② 生成式 Diffusion Model
DiffSplat 支持多种预训练扩散模型作为基础:Stable Diffusion 1.5、SD 2.1、SDXL(1024分辨率)、PixArt-Alpha/Sigma、PixArt-Sigma、Stable Diffusion 3/3.5 以及 FLUX.1。所有这些模型原本都是图像生成器,DiffSplat 通过替换 VAE 解码器(将图像解码器替换为 Gaussian Splat 解码器)并保留冻结的 U-Net/Transformer 主体,把生成目标从"像素值"转变为"Splat 参数"。
支持的 Conditioning 方式包括:
推理时使用 DDIM / DPM-Solver++ / SDE-DPM-Solver++ 等采样器,默认 20 步推理即可获得高质量结果。
③ 轻量级重建模型(GSVAE + GSRecon)用于数据构建
为了构建大规模训练数据,DiffSplat 还训练了一个轻量级的 GSVAE(Gaussian Splatting Variational Autoencoder):Encoder 将多视角 Gaussian Splats 编码到低维 latent 空间,Decoder 则从 latent 重建 Gaussian Splats。这使得可以快速将 GObjaverse 数据集(~26 万个 3D 资产)中的每个对象转换为可学习的 latent 表示,再与多视角渲染图配对构建训练数据集。
此外还有一个专门的 Elevation Estimator(基于 Meta 的 DINOv2 骨干),用于估计每个视角的仰角,辅助构建正确的相机位姿。
项目采用标准的 PyTorch 训练框架,核心依赖包括:
代码组织结构:
src/models/gsrecon.py — GSRecon:核心 Transformer 模型,负责 latent → Gaussian 参数的解码src/models/gsvae.py — GSAutoencoderKL:变分自编码器,支持 TinyAE 加速src/models/gs_render/ — Gaussian 光栅化渲染器(forward/diferred 两种模式)src/models/elevest.py — 仰角估计器(DINOv2 骨干)src/infer_gsdiff_*.py — 推理脚本(按模型架构分版本:SD1.5/SDXL/SD3/PixArt)src/train_gsdiff_*.py — 训练脚本(6000+ 行,包含完整的分布式训练逻辑)extensions/diffusers_diffsplat/ — Diffusers 官方扩展集成(pipeline、transformer、controlnet)推理非常简洁:调用 scripts/infer.sh,传入推理脚本路径 + YAML 配置文件路径 + 模型 tag 即可自动下载权重并运行。推荐使用 HuggingFace 镜像(HF_ENDPOINT=https://hf-mirror.com)加速下载。
显存要求高:虽然推理相对轻量(SD1.5 版约 2GB VRAM),但训练默认配置需要 8×80GB A100,这是普通开发者难以企及的资源门槛,限制了学术复现。
内部数据集依赖:数据集存储路径指向内部 HDFS 目录(<HDFS_DIR>/GObjaverse_parquet),普通用户无法获取完整训练数据,只能下载推理权重做 demo 级别的应用。
无 Web-UI:整个项目是纯 CLI 工具,没有 Gradio/Streamlit 界面,普通用户上手门槛较高。
Camera 固定:推理时相机固定在世界坐标系原点 (0, 0, 1.4),生成后只能做固定视角的旋转展示,尚不支持自由相机轨迹控制。
模型权重需申请:HuggingFace 上的 DiffSplat 预训练权重对某些模型(如 SD3、FLUX)需要额外申请访问权限,不是开箱即用。
DiffSplat 被 ICLR 2025 接收,标志着**"用 2D 生成模型直接做 3D"**这一范式获得了顶级学术会议的认可。与 NeRF 需要数小时优化、LGM 需要多阶段流程相比,DiffSplat 的 end-to-end 1~2 秒推理是一个质的飞跃。
它代表的趋势是:大语言模型和多模态扩散模型中蕴含的海量 2D 视觉先验,可以被巧妙地"劫持"用于 3D 任务,而不必从头爬取或合成昂贵的 3D 数据集。这种"能力迁移"思路与蒸馏(distillation)思想一脉相承,是当前 3D 生成领域最活跃的研究方向之一。
对于开发者而言,DiffSplat 的实用价值在于:作为 3D 资产生成模块嵌入工作流——比如游戏开发中的原型快速迭代、电商平台的产品 3D 化、AR/VR 内容批量生产等。Diffusers 官方扩展的集成意味着未来可以通过 diffusers 库直接调用,降低了工程集成门槛。
# 1. 克隆仓库
git clone https://github.com/chenguolin/DiffSplat.git
cd DiffSplat
# 2. 安装依赖
pip install -r requirements.txt
# 推荐设置 HuggingFace 镜像加速下载
# export HF_ENDPOINT=https://hf-mirror.com
# export HF_HOME=~/.cache/huggingface
# 3. 下载预训练权重
# SD1.5 版本(推荐新手):
huggingface-cli download chenguolin/DiffSplat gsdiff_gobj83k_sd15___render --repo-type model
# 4. 推理(文生 3D)
bash scripts/infer.sh src/infer_gsdiff_sd.py configs/gsdiff_sd15.yaml gsdiff_gobj83k_sd15___render \
--text "a cute robot"
# 5. 推理(图生 3D,需要额外 image_cond 权重)
bash scripts/infer.sh src/infer_gsdiff_sd.py configs/gsdiff_sd15.yaml gsdiff_gobj83k_sd15_image___render \
--image /path/to/your/image.png
生成结果为 3D Gaussian Splats 格式,可使用项目内置的渲染工具导出为可交互的 WebGL 场景,或提取为 PLY/OBJ 文件用于传统 3D 建模软件。