Spaces:
Paused
Paused
| """ | |
| models/text_to_image.py | |
| ─────────────────────── | |
| DiffusersPipelineLoader — wraps any HuggingFace diffusers | |
| StableDiffusionPipeline / StableDiffusionXLPipeline. | |
| Inputs (run kwargs) | |
| ────────────────────────────────────────────────────────────── | |
| prompt : str | |
| negative_prompt : str (optional) | |
| width : int (default from ModelConfig) | |
| height : int (default from ModelConfig) | |
| guidance_scale : float | |
| num_inference_steps : int | |
| seed : int (optional, -1 = random) | |
| Outputs (returned dict) | |
| ────────────────────────────────────────────────────────────── | |
| image : PIL.Image (RGB) | |
| seed : int (seed actually used) | |
| model : str (model_id used) | |
| """ | |
| from __future__ import annotations | |
| import logging | |
| import os | |
| import random | |
| from typing import Any | |
| import torch | |
| from PIL import Image | |
| from models.base_loader import BaseLoader | |
| from utils.device import torch_dtype | |
| logger = logging.getLogger(__name__) | |
| # ─── helpers ────────────────────────────────────────────────────────────────── | |
| def _is_sdxl(model_id: str) -> bool: | |
| """Heuristic: does this model_id look like an SDXL checkpoint?""" | |
| lower = model_id.lower() | |
| return any(k in lower for k in ("xl", "sdxl", "turbo", "lightning")) | |
| def _hf_token() -> str | None: | |
| """Return HF_TOKEN from env if set, else None.""" | |
| return os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN") or None | |
| def _load_pipeline(model_id: str, device: torch.device, dtype: torch.dtype): | |
| """ | |
| Import and return the appropriate diffusers pipeline class, | |
| auto-detecting SD vs SDXL from the model card config. | |
| Strategy | |
| -------- | |
| 1. Try AutoPipelineForText2Image with token (covers auth-gated models). | |
| 2. If that raises an OSError/EnvironmentError related to the Hub (network | |
| unreachable, metadata fetch failure), retry with local_files_only=True | |
| so a previously-cached copy is used instead of hard-failing. | |
| 3. Fall back to plain StableDiffusionPipeline on any other AutoPipeline | |
| failure (e.g. pipeline_tag mismatch on older model cards). | |
| """ | |
| from diffusers import ( | |
| StableDiffusionPipeline, | |
| AutoPipelineForText2Image, | |
| ) | |
| token = _hf_token() | |
| common_kwargs: dict = { | |
| "torch_dtype": dtype, | |
| "safety_checker": None, | |
| "requires_safety_checker": False, | |
| "low_cpu_mem_usage": True, | |
| } | |
| if token: | |
| common_kwargs["token"] = token | |
| def _try_auto(extra: dict | None = None): | |
| kw = {**common_kwargs, **(extra or {})} | |
| return AutoPipelineForText2Image.from_pretrained(model_id, **kw) | |
| def _try_sd(extra: dict | None = None): | |
| kw = {**common_kwargs, **(extra or {})} | |
| return StableDiffusionPipeline.from_pretrained(model_id, **kw) | |
| def _handle_hub_error(err: Exception, attempt: str) -> None: | |
| """ | |
| Translate opaque HF Hub errors into actionable messages, then re-raise. | |
| RepositoryNotFoundError (404/401) | |
| → The model is gated (license agreement required) and no HF_TOKEN | |
| was provided, OR the model_id is wrong. diffusers wraps this as | |
| EnvironmentError with the "not a valid model identifier" text. | |
| Network / cache miss errors | |
| → Hub unreachable; retry with local_files_only handled by caller. | |
| """ | |
| err_str = str(err).lower() | |
| if "not a valid model identifier" in err_str or "repository not found" in err_str: | |
| raise EnvironmentError( | |
| f"\n\n[T2I] Cannot access '{model_id}' on Hugging Face Hub.\n" | |
| f"This model requires you to:\n" | |
| f" 1. Accept its license at https://huggingface.co/{model_id}\n" | |
| f" 2. Generate a HF access token at https://huggingface.co/settings/tokens\n" | |
| f" 3. Add it as HF_TOKEN in your HF Space → Settings → Repository secrets\n" | |
| f"\nOriginal error: {err}" | |
| ) from err | |
| pipe = None | |
| try: | |
| pipe = _try_auto() | |
| except (OSError, EnvironmentError) as hub_err: | |
| err_str = str(hub_err).lower() | |
| # Gated / not found — give actionable message immediately | |
| if "not a valid model identifier" in err_str or "repository not found" in err_str: | |
| _handle_hub_error(hub_err, "AutoPipeline") # always raises | |
| # Network unreachable — retry from local cache | |
| if any(k in err_str for k in ("not cached", "fetch metadata", "connection", "network", | |
| "offline", "cannot reach", "name or service not known")): | |
| logger.warning( | |
| "Hub unreachable for %s. Retrying with local_files_only=True …", model_id, | |
| ) | |
| try: | |
| pipe = _try_auto({"local_files_only": True}) | |
| logger.info("Loaded %s from local cache.", model_id) | |
| except Exception as cache_err: | |
| raise OSError( | |
| f"Cannot load model {model_id}: Hub unreachable and no local cache found.\n" | |
| f" Hub error : {hub_err}\n" | |
| f" Cache error: {cache_err}\n" | |
| "Tip: run the pipeline once with internet access to warm the cache." | |
| ) from cache_err | |
| else: | |
| raise | |
| except Exception as e: | |
| logger.warning("AutoPipeline failed (%s); falling back to StableDiffusionPipeline: %s", model_id, e) | |
| try: | |
| pipe = _try_sd() | |
| except (OSError, EnvironmentError) as hub_err2: | |
| err_str2 = str(hub_err2).lower() | |
| if "not a valid model identifier" in err_str2 or "repository not found" in err_str2: | |
| _handle_hub_error(hub_err2, "StableDiffusionPipeline") # always raises | |
| if any(k in err_str2 for k in ("not cached", "fetch metadata", "connection", "network", | |
| "offline", "cannot reach", "name or service not known")): | |
| logger.warning( | |
| "Hub unreachable for SD fallback %s. Retrying with local_files_only=True …", | |
| model_id, | |
| ) | |
| try: | |
| pipe = _try_sd({"local_files_only": True}) | |
| except Exception as cache_err2: | |
| raise OSError( | |
| f"Cannot load model {model_id}: Hub unreachable and no local cache.\n" | |
| f" Hub error : {hub_err2}\n" | |
| f" Cache error: {cache_err2}\n" | |
| ) from cache_err2 | |
| else: | |
| raise | |
| pipe = pipe.to(device) | |
| # Enable memory-efficient attention when available | |
| try: | |
| pipe.enable_xformers_memory_efficient_attention() | |
| logger.info("xformers attention enabled") | |
| except Exception: | |
| pass | |
| try: | |
| pipe.enable_attention_slicing() | |
| except Exception: | |
| pass | |
| return pipe | |
| # ─── Loader ────────────────────────────────────────────────────────────────── | |
| class DiffusersPipelineLoader(BaseLoader): | |
| def load(self) -> None: | |
| if self._loaded: | |
| logger.info("Already loaded — skipping") | |
| return | |
| logger.info("Loading text-to-image model: %s", self.model_id) | |
| dtype = torch_dtype(self.device) | |
| self.pipe = _load_pipeline(self.model_id, self.device, dtype) | |
| self._loaded = True | |
| logger.info("Model ready: %s", self.model_id) | |
| def run(self, **inputs: Any) -> dict[str, Any]: | |
| if not self._loaded: | |
| self.load() | |
| prompt: str = inputs.get("prompt", "a photograph of a scene") | |
| negative_prompt: str = inputs.get("negative_prompt", "blurry, low quality, distorted") | |
| width: int = int(inputs.get("width", self.kwargs.get("width", 512))) | |
| height: int = int(inputs.get("height", self.kwargs.get("height", 512))) | |
| guidance_scale: float = float(inputs.get("guidance_scale", self.kwargs.get("guidance_scale", 7.5))) | |
| num_steps: int = int(inputs.get("num_inference_steps", self.kwargs.get("num_inference_steps", 25))) | |
| seed: int = int(inputs.get("seed", -1)) | |
| if seed == -1: | |
| seed = random.randint(0, 2**32 - 1) | |
| generator = torch.Generator(device=self.device).manual_seed(seed) | |
| logger.info( | |
| "Generating image prompt=%r size=%dx%d steps=%d cfg=%.1f seed=%d", | |
| prompt[:80], width, height, num_steps, guidance_scale, seed, | |
| ) | |
| call_kwargs: dict[str, Any] = { | |
| "prompt": prompt, | |
| "width": width, | |
| "height": height, | |
| "num_inference_steps": num_steps, | |
| "generator": generator, | |
| } | |
| # guidance_scale=0 is only valid for SDXL-Turbo | |
| if guidance_scale > 0: | |
| call_kwargs["guidance_scale"] = guidance_scale | |
| if negative_prompt: | |
| call_kwargs["negative_prompt"] = negative_prompt | |
| result = self.pipe(**call_kwargs) | |
| image: Image.Image = result.images[0] | |
| return {"image": image, "seed": seed, "model": self.model_id} | |