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] @dataclass 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