sd21_zero / runtime.py
John6666's picture
Upload 4 files
770ad32 verified
Raw
History Blame Contribute Delete
8.22 kB
from __future__ import annotations
import math
import os
import random
import time
import uuid
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Iterable
# ZeroGPU must patch Torch/CUDA before Torch or Diffusers is imported.
import spaces
import torch
from diffusers import EulerDiscreteScheduler, StableDiffusionPipeline
MODEL_IDS = ("Manojb/stable-diffusion-2-1-base",)
DEFAULT_MODEL_ID = MODEL_IDS[0]
MAX_SEED = 2_147_483_647
ZERO_GPU_SIZE = "large"
MIN_GPU_DURATION_SECONDS = 10
MAX_GPU_DURATION_SECONDS = 300
OUTPUT_MAX_AGE_SECONDS = 12 * 60 * 60
OUTPUT_MAX_FILES = 64
OUTPUT_DIR = Path(__file__).resolve().parent / ".runtime" / "outputs"
@dataclass(frozen=True)
class RuntimeInfo:
mode: str
device: str
dtype: torch.dtype
cuda_available: bool
gpu_name: str | None
model_loaded: bool
@property
def is_zerogpu(self) -> bool:
return self.mode == "zerogpu"
@property
def is_assigned_gpu(self) -> bool:
return self.mode == "assigned_gpu"
@property
def is_gpu(self) -> bool:
return self.is_zerogpu or self.is_assigned_gpu
def _detect_runtime() -> RuntimeInfo:
is_zerogpu = bool(os.getenv("SPACES_ZERO_GPU"))
cuda_available = bool(torch.cuda.is_available())
if is_zerogpu:
mode = "zerogpu"
device = "cuda"
dtype = torch.float16
elif cuda_available:
mode = "assigned_gpu"
device = "cuda"
dtype = torch.float16
else:
mode = "cpu"
device = "cpu"
dtype = torch.float32
gpu_name: str | None = None
if cuda_available and not is_zerogpu:
try:
gpu_name = torch.cuda.get_device_name(torch.cuda.current_device())
except Exception:
gpu_name = "CUDA device"
elif is_zerogpu:
gpu_name = f"ZeroGPU {ZERO_GPU_SIZE}"
return RuntimeInfo(
mode=mode,
device=device,
dtype=dtype,
cuda_available=cuda_available,
gpu_name=gpu_name,
model_loaded=False,
)
RUNTIME = _detect_runtime()
def _load_pipelines() -> dict[str, StableDiffusionPipeline]:
pipelines: dict[str, StableDiffusionPipeline] = {}
for model_id in MODEL_IDS:
scheduler = EulerDiscreteScheduler.from_pretrained(
model_id,
subfolder="scheduler",
)
pipeline = StableDiffusionPipeline.from_pretrained(
model_id,
scheduler=scheduler,
torch_dtype=RUNTIME.dtype,
use_safetensors=True,
)
pipeline.to(RUNTIME.device)
pipeline.set_progress_bar_config(disable=False)
pipelines[model_id] = pipeline
return pipelines
if RUNTIME.is_gpu:
torch.set_float32_matmul_precision("high")
PIPELINES = _load_pipelines()
RUNTIME = RuntimeInfo(
mode=RUNTIME.mode,
device=RUNTIME.device,
dtype=RUNTIME.dtype,
cuda_available=RUNTIME.cuda_available,
gpu_name=RUNTIME.gpu_name,
model_loaded=bool(PIPELINES),
)
def startup_summary() -> str:
summary = {
"runtime": RUNTIME.mode,
"device": RUNTIME.device,
"dtype": str(RUNTIME.dtype).replace("torch.", ""),
"gpu": RUNTIME.gpu_name,
"model_loaded": RUNTIME.model_loaded,
"models": list(MODEL_IDS),
}
return f"SD21 runtime: {summary}"
print(startup_summary(), flush=True)
def runtime_banner_markdown() -> str:
if RUNTIME.is_zerogpu:
return (
"**Runtime: ZeroGPU** — the model is preloaded on CUDA emulation and each generation "
f"requests ZeroGPU `{ZERO_GPU_SIZE}` with a workload-based duration."
)
if RUNTIME.is_assigned_gpu:
gpu_name = RUNTIME.gpu_name or "CUDA GPU"
return (
f"**Runtime: assigned GPU** — `{gpu_name}` detected. The pipeline stays resident on "
"CUDA and does not request ZeroGPU quota."
)
return (
"**Runtime: CPU** — the pipeline is loaded in `float32` and generation is available, "
"but inference can be very slow."
)
def require_pipeline(model_id: str) -> StableDiffusionPipeline:
pipeline = PIPELINES.get(model_id)
if pipeline is None:
raise RuntimeError(f"Pipeline not found: {model_id}")
return pipeline
def random_seed() -> int:
return random.SystemRandom().randint(0, MAX_SEED)
def estimate_gpu_duration(
*,
width: int,
height: int,
steps: int,
samples: int,
) -> int:
"""Initial conservative ZeroGPU declaration, calibrated from the prior 10s baseline.
The declaration scales with pixel area, denoising steps, and sequential samples. It is a
maximum allocation request, not a prediction of billed wall time on assigned GPU hardware.
"""
width = max(64, int(width))
height = max(64, int(height))
steps = max(1, int(steps))
samples = max(1, int(samples))
normalized_work = steps * samples * ((width * height) / (512 * 512))
seconds = math.ceil(8.0 + 0.08 * normalized_work)
return max(MIN_GPU_DURATION_SECONDS, min(MAX_GPU_DURATION_SECONDS, seconds))
def estimate_basic_duration(
prompt: str,
negative: str,
scale: float,
model_id: str = DEFAULT_MODEL_ID,
images: Any = None,
*_args: Any,
**_kwargs: Any,
) -> int:
del prompt, negative, scale, model_id, images
return estimate_gpu_duration(width=512, height=512, steps=50, samples=1)
def estimate_advanced_duration(
prompt: str,
negative: str,
scale: float,
width: int,
height: int,
steps: int,
seed: int,
samples: int,
model_id: str = DEFAULT_MODEL_ID,
images: Any = None,
*_args: Any,
**_kwargs: Any,
) -> int:
del prompt, negative, scale, seed, model_id, images
return estimate_gpu_duration(
width=width,
height=height,
steps=steps,
samples=samples,
)
def _gallery_path(item: Any) -> str | None:
if item is None:
return None
if isinstance(item, (str, Path)):
return str(item)
if isinstance(item, dict):
if item.get("path"):
return str(item["path"])
image = item.get("image")
if isinstance(image, dict) and image.get("path"):
return str(image["path"])
if getattr(image, "path", None):
return str(image.path)
return None
if getattr(item, "path", None):
return str(item.path)
image = getattr(item, "image", None)
if getattr(image, "path", None):
return str(image.path)
if isinstance(item, (tuple, list)) and item:
return _gallery_path(item[0])
return None
def _gallery_caption(item: Any) -> str | None:
if isinstance(item, dict):
caption = item.get("caption")
return str(caption) if caption is not None else None
caption = getattr(item, "caption", None)
if caption is not None:
return str(caption)
if isinstance(item, (tuple, list)) and len(item) > 1 and item[1] is not None:
return str(item[1])
return None
def normalize_gallery_items(items: Any) -> list[Any]:
normalized: list[Any] = []
for item in items or []:
path = _gallery_path(item)
if path is None:
continue
caption = _gallery_caption(item)
normalized.append((path, caption) if caption else path)
return normalized
def cleanup_outputs() -> None:
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
now = time.time()
files = sorted(
(path for path in OUTPUT_DIR.glob("*.jpg") if path.is_file()),
key=lambda path: path.stat().st_mtime,
reverse=True,
)
for index, path in enumerate(files):
too_old = now - path.stat().st_mtime > OUTPUT_MAX_AGE_SECONDS
over_limit = index >= OUTPUT_MAX_FILES
if too_old or over_limit:
try:
path.unlink()
except OSError:
pass
def save_image(image: Any, seed: int) -> str:
cleanup_outputs()
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
path = OUTPUT_DIR / f"sd21_{seed}_{uuid.uuid4().hex}.jpg"
image.save(path, format="JPEG", quality=95)
return str(path)
cleanup_outputs()