Bound img2img/inpaint init resolution and exclude the VLM vision tower from NVFP4 text-encoder quant

img2img and inpaint take their output size from the uploaded image and only
snap it to a multiple of 16, so an ordinary phone photo (up to the 4096/side
decode cap, 4x the txt2img 2048 ceiling and ~16x the area) drove an OOM-scale
latent and an opaque 500 on a normal card, while txt2img, upscale, edit, and
FLUX.2-klein inpaint are all already megapixel-bounded. Clamp the init longest
side to 2048 (the txt2img ceiling) before deriving width/height; edit is exempt
since its pipeline resizes to ~1MP internally.

_cast_nvfp4 quantized every nn.Linear with no filter, unlike the int8 and fp8
torchao text-encoder modes which exclude the VLM vision tower / lm_head / T5 wo.
On qwen-image / qwen-image-edit that 4-bit quantized the Qwen2.5-VL image tower,
degrading the edit/image conditioning the sibling schemes protect. Apply the
same make_filter_fn exclusion (require_bf16, mirroring _cast_fp8_dynamic).
This commit is contained in:
Daniel Han 2026-07-08 04:30:24 +00:00
commit 3d27fa7817
4 changed files with 93 additions and 4 deletions

View file

@ -243,6 +243,27 @@ def _snap_to_multiple(img: Any, multiple: int = 16) -> Any:
return img
def _clamp_max_side(img: Any, max_side: int) -> Any:
"""Downscale a PIL image so its longest side is <= ``max_side``, preserving aspect ratio
(high-quality resample); a no-op when it already fits.
img2img / inpaint take their OUTPUT size from the uploaded image, so without a bound an
oversized upload (up to the 4096/side decode cap -- 4x the txt2img 2048 ceiling, ~16x the
area) drives a proportionally larger latent and O(n^2) attention that OOMs the transformer/
VAE on a normal card, surfacing only as an opaque 500. Clamping the longest side to the same
2048 ceiling txt2img enforces (and upscale caps to) keeps these workflows bounded."""
from PIL import Image
w, h = img.size
longest = max(w, h)
if longest <= max_side:
return img
scale = max_side / float(longest)
nw = max(1, int(round(w * scale)))
nh = max(1, int(round(h * scale)))
return img.resize((nw, nh), Image.LANCZOS)
# A small allowlist of well-known official base repos that may load as a full
# (non-GGUF) pipeline even though they are not under ``unsloth/``. These are
# safetensors-only checkpoints from their original publisher (no pickle, no remote
@ -2675,6 +2696,12 @@ class DiffusionBackend:
# txt2img/reference use the validated slider size; upscale already produced a /16
# target. The mask is matched to the snapped image so inpaint stays aligned.
if init_pil is not None and workflow in ("img2img", "inpaint", "edit"):
# img2img/inpaint derive the OUTPUT size from the uploaded image, so bound the
# longest side to txt2img's own 2048 ceiling first -- otherwise a normal phone
# photo (up to the 4096/side decode cap) drives an OOM-scale latent and an
# opaque 500. edit is exempt: its pipeline resizes the input to ~1MP internally.
if workflow in ("img2img", "inpaint"):
init_pil = _clamp_max_side(init_pil, 2048)
init_pil = _snap_to_multiple(init_pil, 16)
if mask_pil is not None and mask_pil.size != init_pil.size:
from PIL import Image as _PILImage

View file

@ -288,9 +288,19 @@ def _cast_fp8(encoder: Any, target: Any) -> None:
def _cast_nvfp4(encoder: Any, target: Any) -> None:
# Weight-only NVFP4: linear weights become 4-bit (packed) NVFP4 tensors and run
# on Blackwell FP4 tensor cores; norms / embeddings (not nn.Linear) are untouched.
# Exclude the VLM vision tower / lm_head / T5 wo and the sub-512 projections, exactly like
# the int8 / fp8 torchao TE modes -- 4-bit-ing a VLM encoder's image tower (qwen-image /
# qwen-image-edit's Qwen2.5-VL) degrades the image/edit conditioning the sibling schemes
# deliberately protect, and require_bf16 skips any non-bf16 Linear the encoder keeps so the
# NVFP4 (scaled_mm-family) cast engages on the bf16 linears instead of aborting the pass.
from torchao.quantization import quantize_
from torchao.prototype.mx_formats import NVFP4WeightOnlyConfig
quantize_(encoder, NVFP4WeightOnlyConfig())
from .diffusion_transformer_quant import DEFAULT_MIN_LINEAR_FEATURES, make_filter_fn
filter_fn = make_filter_fn(
DEFAULT_MIN_LINEAR_FEATURES, _te_exclude_tokens(encoder), require_bf16 = True
)
quantize_(encoder, NVFP4WeightOnlyConfig(), filter_fn = filter_fn)
def _warn(logger: Any, what: str, exc: Exception) -> None:

View file

@ -20,6 +20,7 @@ from core.inference.diffusion import (
DiffusionBackend,
_LoadState,
_base_file_downloaded,
_clamp_max_side,
_resolve_base_repo,
_resolve_diffusion_compute_dtype,
)
@ -42,6 +43,22 @@ from core.inference.diffusion_families import (
# Pure family helpers
def test_clamp_max_side_bounds_oversized_init():
# img2img / inpaint derive the OUTPUT size from the uploaded image; an oversized upload
# (up to the 4096/side decode cap = 4x the txt2img 2048 ceiling) would drive an OOM-scale
# latent. _clamp_max_side bounds the longest side to 2048, preserving aspect ratio.
from PIL import Image
# A 12MP-shaped landscape photo -> longest side clamped to 2048, 4:3 aspect preserved.
out = _clamp_max_side(Image.new("RGB", (4096, 3072)), 2048)
assert out.size == (2048, 1536)
# A portrait upload clamps on its longest (height) side.
assert _clamp_max_side(Image.new("RGB", (1000, 4000)), 2048).size == (512, 2048)
# An image already within bound is returned unchanged (no needless resample).
small = Image.new("RGB", (768, 512))
assert _clamp_max_side(small, 2048) is small
def test_detect_family_from_repo_id():
# Detection is by architecture; Turbo/full and schnell/dev map to one family.
assert detect_family("unsloth/Z-Image-Turbo-GGUF").name == "z-image"

View file

@ -21,6 +21,7 @@ from core.inference.diffusion_precision import (
TE_QUANT_INT8,
TE_QUANT_NVFP4,
_cast_int8_selective,
_cast_nvfp4,
_keep_bf16_block_fqns,
normalize_te_quant,
quantize_text_encoders,
@ -68,13 +69,20 @@ def _stub_casters(monkeypatch, recorder):
hooks.apply_layerwise_casting = lambda module, **kw: recorder.append(("fp8", module))
monkeypatch.setitem(sys.modules, "diffusers.hooks", hooks)
monkeypatch.setitem(sys.modules, "diffusers.hooks.layerwise_casting", casting)
# torchao nvfp4
# torchao nvfp4 -- quantize_ now receives the vision-tower exclusion filter_fn; accept + ignore.
tq = types.ModuleType("torchao.quantization")
tq.quantize_ = lambda module, config: recorder.append(("nvfp4", module))
tq.quantize_ = lambda module, config, filter_fn = None: recorder.append(("nvfp4", module))
mx = types.ModuleType("torchao.prototype.mx_formats")
mx.NVFP4WeightOnlyConfig = lambda: "nvfp4cfg"
monkeypatch.setitem(sys.modules, "torchao.quantization", tq)
monkeypatch.setitem(sys.modules, "torchao.prototype.mx_formats", mx)
# _cast_nvfp4 / _cast_fp8_dynamic pull the shared linear filter from the transformer-quant module.
dtq = types.ModuleType("core.inference.diffusion_transformer_quant")
dtq.DEFAULT_MIN_LINEAR_FEATURES = 512
dtq.make_filter_fn = lambda min_features, exclude = (), *, require_bf16 = False: (
lambda module, fqn = "": True
)
monkeypatch.setitem(sys.modules, "core.inference.diffusion_transformer_quant", dtq)
# ── normalisation ─────────────────────────────────────────────────────────────
@ -308,7 +316,7 @@ def _stub_transformer_quant(monkeypatch, captured):
dtq._make_quant_config = lambda scheme, *a, **k: f"cfg:{scheme}"
dtq.exclude_tokens_for_scheme = lambda scheme: ("modulation",)
def _make_filter_fn(min_features, exclude_name_tokens = ()):
def _make_filter_fn(min_features, exclude_name_tokens = (), *, require_bf16 = False):
def _f(module, fqn = ""):
return not any(tok in fqn for tok in exclude_name_tokens)
@ -329,6 +337,10 @@ def _stub_transformer_quant(monkeypatch, captured):
tq.quantize_ = _quantize_
monkeypatch.setitem(sys.modules, "torchao.quantization", tq)
# _cast_nvfp4 builds its config from here.
mx = types.ModuleType("torchao.prototype.mx_formats")
mx.NVFP4WeightOnlyConfig = lambda: "nvfp4cfg"
monkeypatch.setitem(sys.modules, "torchao.prototype.mx_formats", mx)
def test_int8_filter_keeps_blocks_and_towers_dense(monkeypatch):
@ -353,3 +365,26 @@ def test_int8_filter_keeps_blocks_and_towers_dense(monkeypatch):
assert ff(object(), "visual.blocks.0.attn.qkv") is False
assert ff(object(), "lm_head") is False
assert ff(object(), "model.decoder.wo") is False
def test_nvfp4_filter_keeps_vision_tower_dense(monkeypatch):
# Weight-only NVFP4 on a text encoder must exclude the VLM vision tower / lm_head / T5 "wo"
# like the int8 / fp8 torchao TE modes -- 4-bit-ing a qwen-image(-edit) Qwen2.5-VL image tower
# degrades the edit/image conditioning the sibling schemes deliberately protect. Before the fix
# _cast_nvfp4 quantised every nn.Linear (no filter_fn), so the tower was silently 4-bit.
_stub_torch(monkeypatch)
captured: dict = {}
_stub_transformer_quant(monkeypatch, captured)
enc = types.SimpleNamespace(_keep_in_fp32_modules = ["wo"])
_cast_nvfp4(enc, _target())
assert captured["config"] == "nvfp4cfg"
ff = captured["filter_fn"]
assert ff is not None # a filter is passed now, not None (which quantised everything)
# Vision tower / lm_head / T5 wo stay bf16; an interior projection still quantises.
assert ff(object(), "visual.blocks.0.attn.qkv") is False
assert ff(object(), "vision_tower.encoder.layers.0.mlp.fc1") is False
assert ff(object(), "lm_head") is False
assert ff(object(), "model.decoder.wo") is False
assert ff(object(), "model.layers.5.self_attn.q_proj") is True