Merge remote-tracking branch 'origin/diffusion-sdxl' into diffusion-lora-ux
This commit is contained in:
commit
1437f48e08
8 changed files with 24 additions and 17 deletions
|
|
@ -1538,7 +1538,9 @@ class DiffusionBackend:
|
|||
mask_pil = mask_pil.resize(init_pil.size, _PILImage.NEAREST)
|
||||
if init_pil is not None:
|
||||
# Keep the VAE encode dtype consistent with the input image.
|
||||
self._align_vae_dtype(pipe, getattr(state.family, "denoiser_attr", "transformer"))
|
||||
self._align_vae_dtype(
|
||||
pipe, getattr(state.family, "denoiser_attr", "transformer")
|
||||
)
|
||||
|
||||
# Pipelines vary in which kwargs they accept (img2img derives size from the
|
||||
# input image and may reject width/height; a distilled pipe may take no
|
||||
|
|
|
|||
|
|
@ -225,7 +225,6 @@ def _reset_global_backend_to_native(logger: Any) -> None:
|
|||
AttentionBackendName,
|
||||
_AttentionBackendRegistry,
|
||||
)
|
||||
|
||||
_AttentionBackendRegistry.set_active_backend(AttentionBackendName.NATIVE)
|
||||
except Exception: # noqa: BLE001 — best-effort; leave the global as-is on any change
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -120,7 +120,10 @@ def _cast_fp8(encoder: Any, target: Any) -> None:
|
|||
# gets cast to fp8 and, sharing one tensor, drags the embedding to fp8 with it. The
|
||||
# embedding then emits fp8 activations that crash the first RMSNorm. Skip the tied
|
||||
# projection so the shared tensor stays dense (lm_head is unused for prompt encoding).
|
||||
get_out, get_in = getattr(encoder, "get_output_embeddings", None), getattr(encoder, "get_input_embeddings", None)
|
||||
get_out, get_in = (
|
||||
getattr(encoder, "get_output_embeddings", None),
|
||||
getattr(encoder, "get_input_embeddings", None),
|
||||
)
|
||||
out_emb = get_out() if callable(get_out) else None
|
||||
in_emb = get_in() if callable(get_in) else None
|
||||
if out_emb is not None and in_emb is not None and out_emb.weight is in_emb.weight:
|
||||
|
|
|
|||
|
|
@ -282,6 +282,8 @@ def _enable_cudnn_benchmark(logger: Any) -> bool:
|
|||
except Exception as exc: # noqa: BLE001 — optimisation only
|
||||
_warn(logger, "cudnn_benchmark", exc)
|
||||
return False
|
||||
|
||||
|
||||
# The TF32 flag values from before the first max load flipped them, so a later
|
||||
# non-max load / unload can put the process back exactly as it found it (rather than
|
||||
# forcing a hardcoded default that might clobber another component's choice).
|
||||
|
|
|
|||
|
|
@ -313,6 +313,7 @@ def _make_quant_config(scheme: str, fast_accum: Optional[bool] = None) -> Any:
|
|||
return Float8DynamicActivationFloat8WeightConfig()
|
||||
if scheme == TQ_NVFP4:
|
||||
from torchao.prototype.mx_formats import NVFP4DynamicActivationNVFP4WeightConfig
|
||||
|
||||
# Select the CUTLASS FP4 path, not the default Triton kernel: torchao defaults
|
||||
# use_triton_kernel=True, which needs MSLK installed. On a Blackwell box with the
|
||||
# CUTLASS FP4 extension but no MSLK, the default would make the smoke probe fail
|
||||
|
|
|
|||
|
|
@ -889,9 +889,7 @@ def test_load_sdxl_single_file_uses_pipeline_from_single_file(fake_runtime, tmp_
|
|||
assert status["loaded"] is True
|
||||
assert status["family"] == "sdxl"
|
||||
# The whole-pipeline single-file path was taken with the base repo as config.
|
||||
assert _FakePipeline.last_single_file["path"] == str(
|
||||
(tmp_path / "sdxl.safetensors").resolve()
|
||||
)
|
||||
assert _FakePipeline.last_single_file["path"] == str((tmp_path / "sdxl.safetensors").resolve())
|
||||
assert _FakePipeline.last_single_file["config"] == "stabilityai/stable-diffusion-xl-base-1.0"
|
||||
# The transformer-only single-file build was NOT taken.
|
||||
assert _FakeTransformer.last == {}
|
||||
|
|
|
|||
|
|
@ -487,6 +487,8 @@ def test_invalid_attention_backend_returns_422(client):
|
|||
json = {"model_path": "x/z-image", "gguf_filename": "q.gguf", "attention_backend": "bogus"},
|
||||
)
|
||||
assert resp.status_code == 422
|
||||
|
||||
|
||||
def test_prequant_path_doc_describes_allowlist_not_toggle():
|
||||
# The field help must match the code: UNSLOTH_ALLOW_LOCAL_PREQUANT_PATH is a
|
||||
# directory allowlist, not a =1 toggle (diffusion_prequant._allowed_prequant_roots
|
||||
|
|
|
|||
|
|
@ -46,7 +46,7 @@ def test_sdxl_detection_by_repo_and_override():
|
|||
assert detect_family("stabilityai/sdxl-turbo").name == "sdxl"
|
||||
assert detect_family("some-org/My-Cool-SDXL-Merge").name == "sdxl"
|
||||
assert detect_family("some-org/stable-diffusion-xl-anime").name == "sdxl"
|
||||
assert detect_family("x", override="sdxl").name == "sdxl"
|
||||
assert detect_family("x", override = "sdxl").name == "sdxl"
|
||||
# A GGUF DiT family must NOT be swallowed by the SDXL match.
|
||||
assert detect_family("unsloth/FLUX.1-schnell-GGUF").name == "flux.1"
|
||||
|
||||
|
|
@ -91,9 +91,9 @@ class _FakeVae:
|
|||
self.moved_to = None
|
||||
|
||||
def parameters(self):
|
||||
yield types.SimpleNamespace(dtype=self._dtype)
|
||||
yield types.SimpleNamespace(dtype = self._dtype)
|
||||
|
||||
def to(self, dtype=None):
|
||||
def to(self, dtype = None):
|
||||
self.moved_to = dtype
|
||||
self._dtype = dtype
|
||||
|
||||
|
|
@ -101,29 +101,29 @@ class _FakeVae:
|
|||
def test_align_vae_dtype_uses_unet_denoiser():
|
||||
# For SDXL the denoiser lives at pipe.unet; _align_vae_dtype must read it (a pipe
|
||||
# with only .unet and no .transformer) and cast the VAE to the U-Net's dtype.
|
||||
vae = _FakeVae(dtype="float32")
|
||||
pipe = types.SimpleNamespace(unet=types.SimpleNamespace(dtype="bfloat16"), vae=vae)
|
||||
vae = _FakeVae(dtype = "float32")
|
||||
pipe = types.SimpleNamespace(unet = types.SimpleNamespace(dtype = "bfloat16"), vae = vae)
|
||||
DiffusionBackend._align_vae_dtype(pipe, "unet")
|
||||
assert vae.moved_to == "bfloat16"
|
||||
|
||||
|
||||
def test_align_vae_dtype_transformer_default_unchanged():
|
||||
# DiT default: reads pipe.transformer; a pipe with no transformer is a safe no-op.
|
||||
vae = _FakeVae(dtype="float32")
|
||||
pipe = types.SimpleNamespace(transformer=types.SimpleNamespace(dtype="bfloat16"), vae=vae)
|
||||
vae = _FakeVae(dtype = "float32")
|
||||
pipe = types.SimpleNamespace(transformer = types.SimpleNamespace(dtype = "bfloat16"), vae = vae)
|
||||
DiffusionBackend._align_vae_dtype(pipe)
|
||||
assert vae.moved_to == "bfloat16"
|
||||
# No denoiser attribute -> no-op (does not raise, does not move the VAE).
|
||||
vae2 = _FakeVae(dtype="float32")
|
||||
DiffusionBackend._align_vae_dtype(types.SimpleNamespace(vae=vae2), "unet")
|
||||
vae2 = _FakeVae(dtype = "float32")
|
||||
DiffusionBackend._align_vae_dtype(types.SimpleNamespace(vae = vae2), "unet")
|
||||
assert vae2.moved_to is None
|
||||
|
||||
|
||||
def test_sdxl_lora_supported_on_diffusers():
|
||||
# SDXL is bf16/bnb-4bit on diffusers -> LoRA is allowed (unlike GGUF-via-diffusers).
|
||||
assert diffusion_lora.supports_lora(
|
||||
engine="diffusers", family="sdxl", model_kind="pipeline", transformer_quant=None
|
||||
engine = "diffusers", family = "sdxl", model_kind = "pipeline", transformer_quant = None
|
||||
)
|
||||
assert diffusion_lora.supports_lora(
|
||||
engine="diffusers", family="sdxl", model_kind="single_file", transformer_quant=None
|
||||
engine = "diffusers", family = "sdxl", model_kind = "single_file", transformer_quant = None
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue