本项目实现了一种从卫星图像生成街景图像的跨视角合成方法,基于 Stable Diffusion 潜在扩散模型(Latent Diffusion Model, LDM),核心包括两个训练阶段与两种推理模式。
相关工作参考:
- CS2S (ControlS2S) — 迭代单应性调整(IHA)姿态对齐与文本引导环境控制
- CrossGeo — 跨地理域卫星-街景数据集与泛化方法,作者 Zhongyao Tuo、Qiwei Wang 等
- SPND — 球面全景图像的可变形卷积投影模型
| 项 | 说明 |
|---|---|
| 输入 | RGB 卫星图像(任意分辨率,自动 resize 到 256×256),可选文本提示(如 "in autumn"、"at night") |
| 输出 | 对应位置的 RGB 街景图像(CVUSA 下为 128×512 全景,其他数据集按各自分辨率) |
| 基础工具 | Stable Diffusion v1.4(VAE 编码/解码 + UNet 去噪骨架)、OpenAI CLIP(文本嵌入)、VIT_224(卫星条件编码器) |
两阶段训练,Stage 1 训练 UNet + 卫星条件编码器,Stage 2 冻结基础模型仅微调 SphericalControlNet:
| 阶段 | 训练模块 | 冻结模块 |
|---|---|---|
| Stage 1 | UNet 去噪器 + VIT_224 卫星编码器 + 地面条件模型 | VAE(SD v1.4 权重) |
| Stage 2 | SphericalControlNet(约 388M 参数) | VAE + UNet + VIT_224 |
1. 球面 ControlNet — 全景几何约束
普通卷积对全景图左右边界的拼缝不敏感。本项目的 SphericalControlNet 引入 SphereDeformableConv2d:水平方向使用 F.pad(mode='circular') 实现循环填充,同时用可学习偏移量(offset)自适应调整采样位置,使 ControlNet 能够捕获全景图像的球面几何特性。
2. 动态潜在空间适配
不同数据集产生的潜在表示尺寸不同(如 CVUSA 为 [4,16,64]),硬编码会导致跨数据集崩溃。本项目在运行时从 VAE encoder 输出中动态推导潜在形状(enc.shape),所有采样器、几何变换中的形状参数均由实际编码推导,消除了对特定数据集的依赖。
3. 统一训练入口与 CG 分支修复
原版各数据集使用不同训练脚本,且 Classifier-Free Guidance 分支存在未定义变量的运行时崩溃。本项目以 train_all.py 统一四数据集的两阶段训练,并修复了所有 DDIM 采样器中 CG 分支的引用错误、噪声生成缓存溢出等底层缺陷。
- Python 3.8+
- PyTorch 1.13.1 + CUDA 11.7(>= 12GB 显存)
- PyTorch Lightning 1.9.5
- Stable Diffusion v1.4 checkpoint
conda create -n cs2g python=3.8
conda activate cs2g
conda install pytorch==1.13.1 torchvision==0.14.1 torchaudio==0.13.1 pytorch-cuda=11.7 -c pytorch -c nvidia
pip install pytorch-lightning==1.9.5
pip install omegaconf einops matplotlib opencv-python==4.9.0.80 scikit-image==0.21.0 kornia prefetch_generator lpips pytorch-msssim
pip install -e git+https://github.com/CompVis/taming-transformers.git@master#egg=taming-transformers
pip install -e git+https://github.com/openai/CLIP.git@main#egg=clip
mkdir ckpt && cd ckpt
curl -L https://huggingface.co/CompVis/stable-diffusion-v-1-4-original/resolve/main/sd-v1-4.ckpt -o sd-v1-4.ckpt本项目支持 CVUSA、KITTI、VIGOR 等公开数据集,也可适配其他卫星-街景配对数据集。
| 数据集 | 数据格式 | 来源 |
|---|---|---|
| CVUSA | 卫星 256×256 → 全景 128×512,美国 | Sat2Density |
| KITTI | 卫星 → 前视相机 + 相机内外参,德国 | HighlyAccurate |
| VIGOR | 卫星 → 全景 + GPS/相机参数,4 个美国城市 | VIGOR |
数据集目录结构:
dataset/
├── CVUSA/
├── KITTI_location/
└── VIGOR/
python _verify_crossgeo.py # Stage-1 验证
python _verify_stage2.py # Stage-2 验证以 CVUSA 数据集为例:
python cs2g/train_all.py \
--dataset CVUSA \
--stage 1 \
--devices 0 \
--max-epochs 50 \
--batch-size 2 \
--accumulate-grad-batches 4 \
--lr 1e-5训练产物保存在 outputs/stage1/<时间戳>/checkpoints/,包含 last.ckpt 和按 L1_loss 最优的 epoch checkpoint。
python cs2g/train_all.py \
--dataset CVUSA \
--stage 2 \
--stage1-ckpt outputs/stage1/<run_name>/checkpoints/last.ckpt \
--devices 0 \
--max-epochs 50 \
--batch-size 2 \
--accumulate-grad-batches 4 \
--lr 1e-5训练产物保存在 outputs/stage2/<时间戳>/checkpoints/。
python cs2g/inference_sphere.py \
--config cs2g/configs/Boost_Sat2Den/CVUSA_geo_ldm.yaml \
--ckpt outputs/stage2/<run_name>/checkpoints/last.ckpt \
--sat path/to/satellite_image.png \
--output output.png \
--steps 50 \
--scale 7.5推理参数说明:
| 参数 | 说明 | 默认值 |
|---|---|---|
--config |
YAML 配置文件路径 | 必填 |
--ckpt |
Stage-2 训练得到的 checkpoint | 必填 |
--sat |
输入卫星图像路径(任意分辨率,自动 resize 到 256×256) | 必填 |
--output |
输出街景图像路径 | outputs/output_result.png |
--steps |
DDIM 采样步数 | 50 |
--scale |
CFG 引导强度 | 7.5 |
--device |
运行设备 | cuda |
项目根目录提供 .bat 脚本用于 Windows 环境:
| 脚本 | 功能 |
|---|---|
00_verify.bat |
验证数据、模型、checkpoint 是否就绪 |
01_train_stage1.bat |
Stage-1 训练(支持命令行传参 01_train_stage1.bat <epochs> <batch>) |
02_train_stage2.bat |
Stage-2 训练(自动查找 Stage-1 checkpoint) |
03_inference.bat |
单图推理(自动查找 Stage-2 checkpoint,自动映射 config YAML) |
cs2g/
├── configs/
│ └── Boost_Sat2Den/
│ ├── crossgeo_geo_ldm_stage1.yaml # CrossGeo Stage-1 配置
│ ├── crossgeo_geo_ldm_sphere.yaml # CrossGeo Stage-2 配置(含 SphericalControlNet)
│ ├── CVUSA_geo_ldm.yaml # CVUSA 配置
│ ├── KITTI_geo_ldm.yaml # KITTI 配置
│ └── VIGOR_geo_ldm.yaml # VIGOR 配置
├── models/
│ ├── crossgeo_geo_ldm/
│ │ └── crossgeo_txt_control.py # ★ CrossGEO_Sat2Den_ddpm 主模型(LightningModule)
│ ├── crossgeo_geo_ldm_diffusion/
│ │ ├── crossgeo_latent_diffusion.py # CrossGeo DDPM(扩散过程 + p_losses)
│ │ ├── crossgeo_ddim.py # DDIM 采样器(含 SphericalControlNet 分支)
│ │ └── openaimodel.py # UNet 去噪器
│ ├── CVUSA_geo_ldm/ # CVUSA 条件模型
│ ├── CVUSA_geo_ldm_diffusion/ # CVUSA 扩散模型 + DDIM 采样器
│ ├── KITTI_geo_ldm/ # KITTI 条件模型
│ ├── KITTI_geo_ldm_diffusion/ # KITTI 扩散模型 + DDIM 采样器
│ ├── VIGOR_geo_ldm/ # VIGOR 条件模型
│ ├── VIGOR_geo_ldm_diffusion/ # VIGOR 扩散模型 + DDIM 采样器
│ ├── spherical_controlnet.py # ★ SphericalControlNet (SphereDeformableConv2d)
│ ├── geometry/ # 几何变换 (sat2grd, grd2sat, 投影映射)
│ ├── loss_fun/ # 损失函数
│ ├── eval/ # 评估指标
│ └── autoencoder/ # VAE 自编码器
├── ldm/ # 潜在扩散模型基础库
│ ├── modules/ # CrossAttention, SpatialTransformer, 编码器
│ └── models/diffusion/ # DDIM/DDPM 基类
├── dataloader/
│ ├── crossgeo_txt.py # CrossGeo 数据集
│ ├── CVUSA_txt.py # CVUSA 数据集
│ ├── KITTI_wo_loc.py # KITTI 数据集
│ └── VIGOR_corr.py # VIGOR 数据集
├── utils/ # 工具函数 (instantiate_from_config, callback 等)
├── train_all.py # ★ 统一训练入口
├── visualization.py # 批量推理/可视化入口
├── inference_sphere.py # SphericalControlNet 单图推理入口
└── main.py # 原始训练入口
| 参数 | 说明 | 默认值 |
|---|---|---|
--dataset |
数据集: CVUSA / KITTI / VIGOR(可扩展) |
CVUSA |
--stage |
训练阶段: 1(基础)/ 2(SphericalControlNet) |
1 |
--devices |
GPU 设备号 | "0" |
--max-epochs |
最大训练轮数 | 50 |
--batch-size |
批次大小 | 2 |
--accumulate-grad-batches |
梯度累积步数 | 4 |
--lr |
学习率(覆盖 YAML 配置中的值) | None(使用配置文件值) |
--stage1-ckpt |
Stage-2 训练时加载的 Stage-1 checkpoint | None |
--resume |
恢复训练的 checkpoint 路径(仅 Stage-1) | None |
--seed |
随机种子 | 24 |
如有问题,欢迎联系:couplechein@163.com
