Instructions to use BiliSakura/DiffusionSat-Single-512 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use BiliSakura/DiffusionSat-Single-512 with Diffusers:
pip install -U diffusers transformers accelerate
from diffusers import ControlNetModel, StableDiffusionControlNetPipeline controlnet = ControlNetModel.from_pretrained("BiliSakura/DiffusionSat-Single-512") pipe = StableDiffusionControlNetPipeline.from_pretrained( "fill-in-base-model", controlnet=controlnet ) - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- Draw Things
- DiffusionBee
Download run_demo_inference.py from BiliSakura/DiffusionSat-Single-512: direct link, hf CLI and curl.
- Browser
- Download file 5.26 kB
-
https://huggingface.co/BiliSakura/DiffusionSat-Single-512/resolve/main/run_demo_inference.py
- Command line
-
hf download hf://BiliSakura/DiffusionSat-Single-512/run_demo_inference.py
-
curl -L -o run_demo_inference.py https://huggingface.co/BiliSakura/DiffusionSat-Single-512/resolve/main/run_demo_inference.py
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() | |