# Conflicts: # scripts/diffusion_bench.py # scripts/diffusion_quality.py # studio/backend/core/inference/diffusion.py # studio/backend/core/inference/diffusion_device.py # studio/backend/core/inference/diffusion_families.py # studio/backend/core/inference/diffusion_memory.py # studio/backend/core/inference/diffusion_precision.py # studio/backend/core/inference/diffusion_speed.py # studio/backend/models/inference.py # studio/backend/routes/inference.py # studio/backend/tests/test_diffusion_backend.py # studio/backend/tests/test_diffusion_device.py # studio/backend/tests/test_diffusion_memory.py # studio/backend/tests/test_diffusion_precision.py # studio/backend/tests/test_diffusion_speed.py
152 lines
6.7 KiB
Python
152 lines
6.7 KiB
Python
# SPDX-License-Identifier: AGPL-3.0-only
|
|
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
|
|
|
"""Pure helpers for diffusion model identification.
|
|
|
|
No torch/diffusers imports here: everything in this module is a pure function of
|
|
its string/path arguments so it can be unit-tested without the heavy runtime.
|
|
|
|
A diffusion checkpoint published as a single-file GGUF only carries the
|
|
transformer weights; the matching VAE / text encoders / scheduler come from a
|
|
companion ``diffusers`` base repo. ``DiffusionFamily`` maps a checkpoint to the
|
|
diffusers classes and base repo needed to assemble the full pipeline.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import re
|
|
from dataclasses import dataclass, field
|
|
from pathlib import Path, PurePosixPath
|
|
from typing import Optional
|
|
|
|
|
|
@dataclass(frozen = True)
|
|
class DiffusionFamily:
|
|
name: str
|
|
pipeline_class: str
|
|
transformer_class: str
|
|
base_repo: str
|
|
# Pipeline kwarg carrying the guidance value. Most use "guidance_scale";
|
|
# Qwen-Image's distilled guidance is off, so its real CFG is "true_cfg_scale".
|
|
cfg_kwarg: str = "guidance_scale"
|
|
# Extra lowercased substrings (besides ``name``) that map a repo id here.
|
|
aliases: tuple[str, ...] = field(default_factory = tuple)
|
|
# True for families whose activations overflow float16's finite range
|
|
# (~6.5e4) and produce inf -> NaN latents -> a black image. The backend
|
|
# promotes a resolved float16 to float32 for these at load time.
|
|
fp16_incompatible: bool = False
|
|
# False for families whose denoiser block doesn't compile cleanly with
|
|
# regional torch.compile (Z-Image). Only consulted on the non-GGUF path; the
|
|
# GGUF transformer is never compiled regardless.
|
|
supports_torch_compile: bool = True
|
|
|
|
|
|
# Keyed by architecture, not per model variant: a checkpoint's specific base repo
|
|
# is read from its HF base_model tag at load time, so one entry covers Turbo/full,
|
|
# schnell/dev, etc. base_repo here is only a fallback. Only archs whose diffusers
|
|
# transformer supports from_single_file load here (ERNIE-Image does not, yet).
|
|
# FLUX.2-dev and FLUX.2-klein-9B are left out only because their base diffusers
|
|
# repos are gated; the open klein-4B base stands in for the klein family below.
|
|
_FAMILIES: tuple[DiffusionFamily, ...] = (
|
|
DiffusionFamily(
|
|
name = "flux.1",
|
|
pipeline_class = "FluxPipeline",
|
|
transformer_class = "FluxTransformer2DModel",
|
|
base_repo = "black-forest-labs/FLUX.1-schnell",
|
|
aliases = ("flux1", "flux-1"),
|
|
),
|
|
# FLUX.2-klein is a distinct pipeline (Flux2KleinPipeline) with a Qwen3 text
|
|
# encoder, not the Mistral-based Flux2Pipeline; it must precede a generic
|
|
# flux match. The base Flux2Pipeline (FLUX.2-dev) is gated, so it's omitted.
|
|
DiffusionFamily(
|
|
name = "flux.2-klein",
|
|
pipeline_class = "Flux2KleinPipeline",
|
|
transformer_class = "Flux2Transformer2DModel",
|
|
base_repo = "black-forest-labs/FLUX.2-klein-4B",
|
|
aliases = ("flux2-klein",),
|
|
),
|
|
DiffusionFamily(
|
|
name = "qwen-image",
|
|
pipeline_class = "QwenImagePipeline",
|
|
transformer_class = "QwenImageTransformer2DModel",
|
|
base_repo = "Qwen/Qwen-Image",
|
|
cfg_kwarg = "true_cfg_scale",
|
|
aliases = ("qwen_image", "qwenimage"),
|
|
),
|
|
DiffusionFamily(
|
|
name = "z-image",
|
|
pipeline_class = "ZImagePipeline",
|
|
transformer_class = "ZImageTransformer2DModel",
|
|
base_repo = "Tongyi-MAI/Z-Image-Turbo",
|
|
aliases = ("zimage", "z_image"),
|
|
# Z-Image's MLP down-projections peak near 9e5, which overflows float16.
|
|
fp16_incompatible = True,
|
|
# Z-Image's denoiser block is excluded from regional torch.compile.
|
|
supports_torch_compile = False,
|
|
),
|
|
)
|
|
|
|
# Editing / inpaint checkpoints share an arch keyword but need a different
|
|
# pipeline and an input image, which this text-to-image backend doesn't drive.
|
|
_EDIT_KEYWORDS = ("edit", "kontext", "inpaint", "inpainting")
|
|
|
|
|
|
def detect_family(repo_id: str, override: Optional[str] = None) -> Optional[DiffusionFamily]:
|
|
"""Resolve a ``DiffusionFamily`` from a repo id, or an explicit override.
|
|
|
|
``override`` matches a family ``name`` or alias exactly; otherwise the repo
|
|
id is scanned for the first family whose name/alias appears in it. Image
|
|
editing checkpoints are rejected (None) since this backend is text-to-image.
|
|
"""
|
|
if override:
|
|
key = override.strip().lower()
|
|
for fam in _FAMILIES:
|
|
if key == fam.name or key in fam.aliases:
|
|
return fam
|
|
return None
|
|
needle = repo_id.lower()
|
|
# Match edit keywords as whole id segments, not raw substrings, so a normal
|
|
# text-to-image repo like ".../some-image-edition" isn't misread as an editing
|
|
# checkpoint. Qwen-Image-Edit / FLUX.1-Kontext still match (edit/kontext are
|
|
# whole tokens there). Split on both path separators so a Windows local path
|
|
# is segmented too.
|
|
segments = set(re.split(r"[-_./\\]+", needle))
|
|
if any(kw in segments for kw in _EDIT_KEYWORDS):
|
|
return None
|
|
for fam in _FAMILIES:
|
|
if fam.name in needle or any(alias in needle for alias in fam.aliases):
|
|
return fam
|
|
return None
|
|
|
|
|
|
def resolve_base_repo(fam: DiffusionFamily, base_repo: Optional[str]) -> str:
|
|
"""The companion diffusers repo: caller-supplied if given, else the family fallback."""
|
|
base = (base_repo or "").strip()
|
|
return base or fam.base_repo
|
|
|
|
|
|
def resolve_local_gguf_child(repo_root: Path, gguf_filename: str) -> Path:
|
|
"""Resolve ``gguf_filename`` to a file under ``repo_root``, rejecting escapes.
|
|
|
|
``gguf_filename`` is user-supplied, so an absolute path or a ``..`` segment
|
|
could otherwise read outside the local repo directory.
|
|
"""
|
|
if (
|
|
Path(gguf_filename).is_absolute()
|
|
or PurePosixPath(gguf_filename).is_absolute()
|
|
or gguf_filename.startswith(("/", "\\"))
|
|
or "\\" in gguf_filename
|
|
):
|
|
raise ValueError("gguf_filename must be a relative path inside the repo.")
|
|
rel = PurePosixPath(gguf_filename)
|
|
if any(part in ("", ".", "..") for part in rel.parts):
|
|
raise ValueError("gguf_filename must not contain '', '.', or '..' segments.")
|
|
# Resolve symlinks before the containment check: the guards above stop
|
|
# lexical escapes, but a symlink inside the repo could still point outside it.
|
|
repo_real = repo_root.resolve()
|
|
child = repo_root.joinpath(*rel.parts).resolve()
|
|
if child != repo_real and repo_real not in child.parents:
|
|
raise ValueError("gguf_filename must resolve to a file inside the repo.")
|
|
if not child.is_file():
|
|
raise FileNotFoundError(f"'{gguf_filename}' is not a file under {repo_root}.")
|
|
return child
|