#!/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()