DiffusionSat-Single-512 / run_demo_inference.py
BiliSakura's picture
Add files using upload-large-folder tool
28f2247 verified
Raw History Blame Contribute Delete
5.26 kB
#!/usr/bin/env python3
"""Run DiffusionSat-Single-512 demo inference and save to demo_images/."""
import argparse
from pathlib import Path
from typing import Optional
import torch
from diffusers import (
AutoencoderKL,
DDIMScheduler,
DPMSolverMultistepScheduler,
EulerDiscreteScheduler,
PNDMScheduler,
)
from transformers import CLIPTokenizer
try:
from transformers import CLIPImageProcessor, CLIPTextModel
except ImportError:
from transformers import CLIPFeatureExtractor as CLIPImageProcessor # type: ignore
from transformers import CLIPTextModel # type: ignore
from pipeline_diffusionsat import DiffusionSatPipeline
from unet.sat_unet import SatUNet
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Run DiffusionSat demo inference.")
parser.add_argument(
"--scheduler",
choices=["ddim", "pndm", "euler", "dpmpp_2m"],
default="ddim",
help="Scheduler to use for sampling.",
)
parser.add_argument(
"--prompt",
default="satellite image of farmland with roads",
help="Text prompt for generation.",
)
parser.add_argument("--steps", type=int, default=50, help="Number of inference steps.")
parser.add_argument("--guidance-scale", type=float, default=7.5, help="Classifier-free guidance scale.")
parser.add_argument("--seed", type=int, default=42, help="Random seed.")
parser.add_argument(
"--metadata",
choices=["none", "zeros", "custom"],
default="zeros",
help="Metadata mode: none/zeros/custom.",
)
parser.add_argument(
"--metadata-values",
default=None,
help=(
"Comma-separated normalized metadata values for custom mode: "
"lon,lat,gsd,cloud,year,month,day"
),
)
parser.add_argument(
"--output",
default="output.jpeg",
help="Output filename under demo_images/.",
)
return parser.parse_args()
def build_scheduler(name: str, config) -> object:
if name == "ddim":
return DDIMScheduler.from_config(config)
if name == "pndm":
return PNDMScheduler.from_config(config)
if name == "euler":
return EulerDiscreteScheduler.from_config(config)
if name == "dpmpp_2m":
return DPMSolverMultistepScheduler.from_config(config)
raise ValueError(f"Unsupported scheduler: {name}")
def parse_and_validate_metadata(args: argparse.Namespace) -> Optional[list[float]]:
if args.metadata == "none":
return None
if args.metadata == "zeros":
return [0.0] * 7
if not args.metadata_values:
raise ValueError("--metadata-values is required when --metadata custom.")
parts = [part.strip() for part in args.metadata_values.split(",")]
if len(parts) != 7:
raise ValueError("Expected exactly 7 metadata values: lon,lat,gsd,cloud,year,month,day")
try:
values = [float(part) for part in parts]
except ValueError as exc:
raise ValueError("All metadata values must be numeric floats.") from exc
field_names = ["lon", "lat", "gsd", "cloud", "year", "month", "day"]
# Training code normalizes metadata to roughly [0, 1000]. We enforce
# this schema to avoid silent out-of-distribution conditioning.
for field_name, value in zip(field_names, values):
if not torch.isfinite(torch.tensor(value)):
raise ValueError(f"Metadata field '{field_name}' is not finite: {value}")
if value < 0.0 or value > 1000.0:
raise ValueError(
f"Metadata field '{field_name}'={value} is outside expected normalized range [0, 1000]."
)
return values
def main() -> None:
args = parse_args()
repo = Path(__file__).resolve().parent
demo_dir = repo / "demo_images"
demo_dir.mkdir(exist_ok=True)
vae = AutoencoderKL.from_pretrained(str(repo / "vae"), local_files_only=True)
text_encoder = CLIPTextModel.from_pretrained(str(repo / "text_encoder"), local_files_only=True)
tokenizer = CLIPTokenizer.from_pretrained(str(repo / "tokenizer"), local_files_only=True)
unet = SatUNet.from_pretrained(str(repo / "unet"), local_files_only=True)
scheduler = DDIMScheduler.from_pretrained(str(repo / "scheduler"), local_files_only=True)
feature_extractor = CLIPImageProcessor.from_pretrained(
str(repo / "feature_extractor"), local_files_only=True
)
pipe = DiffusionSatPipeline(
vae=vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
unet=unet,
scheduler=scheduler,
safety_checker=None,
feature_extractor=feature_extractor,
requires_safety_checker=False,
).to("cuda")
pipe.scheduler = build_scheduler(args.scheduler, pipe.scheduler.config)
metadata = parse_and_validate_metadata(args)
output_path = demo_dir / args.output
image = pipe(
prompt=args.prompt,
metadata=metadata,
num_inference_steps=args.steps,
guidance_scale=args.guidance_scale,
generator=torch.Generator(device="cuda").manual_seed(args.seed),
).images[0]
image.save(output_path)
print(f"Saved {output_path}")
if __name__ == "__main__":
main()