unsloth/studio/backend/core/inference/diffusion_families.py
Daniel Han 48628252bd Merge remote-tracking branch 'origin/image-generation' into diffusion-phase4-native
# 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
2026-07-01 11:19:15 +00:00

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