Download src/generate.py from Nate-M/StyleX: direct link, hf CLI and curl.
- Browser
- Download file 4.79 kB
-
https://huggingface.co/spaces/Nate-M/StyleX/resolve/main/src/generate.py
- Command line
-
hf download hf://spaces/Nate-M/StyleX/src/generate.py
-
curl -L -o generate.py https://huggingface.co/spaces/Nate-M/StyleX/resolve/main/src/generate.py
4.79 kB
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| from typing import List, Optional | |
| from PIL import Image | |
| from .styles import Style | |
| from .utils import timestamp | |
| IMG_EXTS = {".png", ".jpg", ".jpeg", ".webp"} | |
| def _list_images(folder: Path) -> List[Path]: | |
| if not folder.exists(): | |
| return [] | |
| files: List[Path] = [] | |
| for p in folder.iterdir(): | |
| if p.is_file() and p.suffix.lower() in IMG_EXTS: | |
| files.append(p) | |
| return sorted(files) | |
| def _load_pil(path: Path) -> Image.Image: | |
| return Image.open(path).convert("RGB") | |
| def _pick_reference_images( | |
| style_folder: Path, | |
| max_images: int, | |
| sample_mode: str = "random", | |
| seed: Optional[int] = None, | |
| ) -> List[Image.Image]: | |
| import torch | |
| paths = _list_images(style_folder) | |
| if not paths or max_images <= 0: | |
| return [] | |
| max_images = min(max_images, len(paths)) | |
| if sample_mode == "first": | |
| chosen = paths[:max_images] | |
| else: | |
| generator = torch.Generator(device="cpu") | |
| generator.manual_seed(0 if seed is None else int(seed)) | |
| idx = torch.randperm(len(paths), generator=generator).tolist() | |
| chosen = [paths[i] for i in idx[:max_images]] | |
| return [_load_pil(p) for p in chosen] | |
| class GenerateConfig: | |
| model_id: str | |
| steps: int = 6 | |
| guidance: float = 1.0 | |
| height: int = 512 | |
| width: int = 512 | |
| device: str = "cuda" | |
| torch_dtype_name: str = "bfloat16" | |
| cpu_offload: bool = True | |
| use_ref_images: bool = True | |
| ref_max_images: int = 4 | |
| ref_sample_mode: str = "random" | |
| seed: Optional[int] = None | |
| num_outputs: int = 1 | |
| def _is_flux2(model_id: str) -> bool: | |
| mid = model_id.lower() | |
| return "flux.2" in mid or "flux2" in mid | |
| def _resolve_torch_dtype(dtype_name: str): | |
| import torch | |
| if dtype_name == "float16": | |
| return torch.float16 | |
| if dtype_name == "float32": | |
| return torch.float32 | |
| return torch.bfloat16 | |
| def _load_flux2_pipeline(model_id: str, cfg: GenerateConfig): | |
| import torch | |
| from diffusers import Flux2KleinPipeline, Flux2Pipeline | |
| torch_dtype = _resolve_torch_dtype(cfg.torch_dtype_name) | |
| if "klein" in model_id.lower(): | |
| pipe = Flux2KleinPipeline.from_pretrained(model_id, torch_dtype=torch_dtype) | |
| else: | |
| pipe = Flux2Pipeline.from_pretrained(model_id, torch_dtype=torch_dtype) | |
| if cfg.device == "cuda" and torch.cuda.is_available(): | |
| if cfg.cpu_offload: | |
| pipe.enable_model_cpu_offload() | |
| else: | |
| pipe.to("cuda") | |
| else: | |
| pipe.to("cpu") | |
| return pipe | |
| def generate_one( | |
| user_prompt: str, | |
| style: Style, | |
| out_root: Path, | |
| cfg: GenerateConfig, | |
| ) -> List[Path]: | |
| import torch | |
| out_root.mkdir(parents=True, exist_ok=True) | |
| if not _is_flux2(cfg.model_id): | |
| raise ValueError( | |
| "This generate.py expects a FLUX.2 model_id " | |
| f"(got: {cfg.model_id})" | |
| ) | |
| pipe = _load_flux2_pipeline(cfg.model_id, cfg) | |
| if cfg.seed is None: | |
| generator = None | |
| else: | |
| gen_device = "cuda" if (cfg.device == "cuda" and torch.cuda.is_available()) else "cpu" | |
| generator = torch.Generator(device=gen_device).manual_seed(int(cfg.seed)) | |
| ref_images: Optional[List[Image.Image]] = None | |
| if cfg.use_ref_images: | |
| ref_images = _pick_reference_images( | |
| style.folder, | |
| max_images=cfg.ref_max_images, | |
| sample_mode=cfg.ref_sample_mode, | |
| seed=cfg.seed, | |
| ) | |
| if ref_images: | |
| print( | |
| f"FLUX.2 reference images used: {len(ref_images)} " | |
| f"(mode={cfg.ref_sample_mode}, max={cfg.ref_max_images})" | |
| ) | |
| else: | |
| print("FLUX.2 reference images: none found (running txt2img).") | |
| call_kwargs = dict( | |
| prompt=user_prompt, | |
| height=int(cfg.height), | |
| width=int(cfg.width), | |
| num_inference_steps=max(1, int(cfg.steps)), | |
| guidance_scale=float(cfg.guidance), | |
| generator=generator, | |
| num_images_per_prompt=max(1, int(cfg.num_outputs)), | |
| ) | |
| if ref_images: | |
| call_kwargs["image"] = ref_images | |
| result = pipe(**call_kwargs) | |
| out_images: List[Image.Image] = list(result.images) | |
| time_id = timestamp() | |
| out_dir = out_root / style.name | |
| out_dir.mkdir(parents=True, exist_ok=True) | |
| paths: List[Path] = [] | |
| for i, img in enumerate(out_images, start=1): | |
| out_path = out_dir / f"{time_id}_{i:02d}.png" | |
| img.save(out_path) | |
| paths.append(out_path) | |
| return paths |