Views
No views yet
[!NOTE] If you encounter pipeline loading failure or unexpected output, please contact bili_sakura@zju.edu.cn.
diffusers.DiffusionPipeline.from_pretrained().model_index.json is set to the default text-to-image pipeline (DiffusionSatPipeline) so DiffusionPipeline.from_pretrained() works out of the box. The ControlNet variant is loaded via custom_pipeline plus the controlnet subfolder, as shown below.pipeline_diffusionsat.py: Standard text-to-image pipeline with DiffusionSat metadata support.pipeline_diffusionsat_controlnet.py: ControlNet pipeline with DiffusionSat metadata and conditional metadata support.ckpt/diffusionsat/) should contain the standard diffusers components (unet, vae, scheduler, etc.). You can reference these pipeline files directly from this directory or copy them to your checkpoint folder.pipeline_diffusionsat.py for standard generation.1import torch
2from diffusers import DiffusionPipeline
3
4# Load pipeline
5pipe = DiffusionPipeline.from_pretrained(
6 "path/to/ckpt/diffusionsat",
7 custom_pipeline="./pipeline_diffusionsat.py", # Path to this file
8 torch_dtype=torch.float16,
9 trust_remote_code=True,
10)
11pipe = pipe.to("cuda")
12
13# Optional: Metadata (normalized lat, lon, timestamp, GSD, etc.)
14# metadata = [0.5, -0.3, 0.7, 0.2, 0.1, 0.0, 0.5]
15
16# Generate
17image = pipe(
18 "satellite image of farmland",
19 metadata=None, # Optional
20 height=512,
21 width=512,
22 num_inference_steps=30,
23).images[0]pipeline_diffusionsat_controlnet.py for ControlNet generation.1import torch
2import numpy as np
3from PIL import Image
4from diffusers import DiffusionPipeline, ControlNetModel
5
6# 1. Load the fMoW 2D ControlNet
7controlnet = ControlNetModel.from_pretrained(
8 "path/to/ckpt/diffusionsat/controlnet",
9 torch_dtype=torch.float16,
10 conditioning_channels=10,
11)
12
13# 2. Load pipeline with ControlNet
14pipe = DiffusionPipeline.from_pretrained(
15 "path/to/ckpt/diffusionsat",
16 controlnet=controlnet,
17 custom_pipeline="./pipeline_diffusionsat_controlnet.py", # Path to this file
18 torch_dtype=torch.float16,
19 trust_remote_code=True,
20)
21pipe = pipe.to("cuda")
22
23# 3. Prepare 10-channel conditioning tensor (RGB + 7 zero channels)
24control_image = Image.open("path/to/conditioning_image.png").convert("RGB").resize((256, 256))
25rgb = torch.from_numpy(np.array(control_image)).permute(2, 0, 1).unsqueeze(0).float() / 255.0
26extra = torch.zeros((1, 7, 256, 256), dtype=rgb.dtype)
27control_tensor = torch.cat([rgb, extra], dim=1).to(device="cuda", dtype=torch.float16)
28
29# 4. Generate
30image = pipe(
31 prompt="satellite image of farmland",
32 image=control_tensor,
33 metadata=None,
34 cond_metadata=None,
35 height=256,
36 width=256,
37 num_inference_steps=30,
38).images[0]