Wire the hosted pre-cast fp8 text encoders

qwen-image and flux.2-dev (diffusion) and ltx-2 (video) now resolve a
hosted pre-cast fp8 text encoder from their unsloth -FP8 repos:

- unsloth/Qwen-Image-FP8: Qwen2.5-VL-7B, 16.6 GB dense -> 8.8 GB
- unsloth/FLUX.2-dev-FP8: Mistral-Small-24B, 48.0 GB dense -> 24.7 GB
- unsloth/LTX-2-FP8: Gemma3-12B, 48.7 GB fp32 store -> 13.2 GB

Every checkpoint verified bit-identical to dense-load-then-cast
(729 / 585 / 1066 tensors, zero mismatches) and smoke-tested through the
real backends with the repo engagement marker. Tests cover the wired
entries, the resolver filenames, builder metadata weights_only survival,
and the idempotent re-cast.
This commit is contained in:
Daniel Han 2026-07-18 06:52:46 +00:00
commit 99486d8f6e
3 changed files with 91 additions and 2 deletions

View file

@ -196,6 +196,8 @@ _FAMILIES: tuple[DiffusionFamily, ...] = (
("int8", "unsloth/FLUX.2-dev-FP8"),
("fp8", "unsloth/FLUX.2-dev-FP8"),
),
# Pre-cast Mistral-Small-24B conditioner (bf16 ~48 GB dense, ~24.7 GB pre-cast).
te_prequant_repos = (("fp8", "text_encoder", "unsloth/FLUX.2-dev-FP8"),),
aliases = ("flux2-dev", "flux2dev"),
# LoRA training via the DiT trainer (QLoRA nf4 by default); the base repo is gated, so
# training requires an HF token with the FLUX.2-dev license accepted.
@ -247,6 +249,9 @@ _FAMILIES: tuple[DiffusionFamily, ...] = (
base_repo = "Qwen/Qwen-Image",
# int8 only: fp8 is family-denied (_FAMILY_SCHEME_DENY) so a repo entry would be dead.
prequant_repos = (("int8", "unsloth/Qwen-Image-FP8"),),
# Pre-cast Qwen2.5-VL-7B (bf16 ~16.6 GB dense, ~8.8 GB pre-cast). The DiT fp8 denial
# is a transformer-scheme rule; the layerwise TE cast is unaffected.
te_prequant_repos = (("fp8", "text_encoder", "unsloth/Qwen-Image-FP8"),),
cfg_kwarg = "true_cfg_scale",
aliases = ("qwen_image", "qwenimage"),
# LoRA training via the DiT trainer, defaulting to the prequant nf4 repo (QLoRA).

View file

@ -100,6 +100,9 @@ _FAMILIES: tuple[VideoFamily, ...] = (
# transformer 37.8 bf16; Gemma3-27B TE ~50.4; VAE 2.4 + connectors 2.9 + audio 0.2.
bf16_components_gb = (37.8, 50.4, 5.5),
gguf_repo = "unsloth/LTX-2.3-GGUF",
# Pre-cast Gemma3-12B TE (hub store is fp32 ~49 GB, pre-cast ~13.2 GB): the biggest
# download win of the hosted TE set.
te_prequant_repos = (("fp8", "text_encoder", "unsloth/LTX-2-FP8"),),
),
# Wan2.2-TI2V-5B (diffusers >= 0.35, verified on 0.39): ~5B single-stream video DiT (UMT5
# text encoder). No audio, no second expert (boundary_ratio null, transformer_2 null), so

View file

@ -10,6 +10,7 @@ gating -- all without CUDA, the Hub, or a real transformers model."""
from __future__ import annotations
import types
from pathlib import Path
import pytest
@ -231,7 +232,87 @@ def test_family_dataclasses_declare_te_prequant_field():
assert DiffusionFamily.__dataclass_fields__["te_prequant_repos"].default_factory is tuple
assert VideoFamily.__dataclass_fields__["te_prequant_repos"].default_factory is tuple
# No family ships a hosted TE checkpoint yet: the campaign wires entries after the
# artifacts are gate-validated and uploaded.
# Families without a hosted TE checkpoint keep the empty default.
fam = detect_family("unsloth/FLUX.1-schnell-GGUF")
assert fam.te_prequant_repos == ()
def test_hosted_te_prequant_entries():
"""The hosted pre-cast fp8 text encoders live in the family's own -FP8 repos."""
from core.inference.diffusion_families import detect_family
from core.inference.video_families import detect_video_family
assert detect_family("Qwen/Qwen-Image").te_prequant_repos == (
("fp8", "text_encoder", "unsloth/Qwen-Image-FP8"),
)
assert detect_family("black-forest-labs/FLUX.2-dev").te_prequant_repos == (
("fp8", "text_encoder", "unsloth/FLUX.2-dev-FP8"),
)
assert detect_video_family("Lightricks/LTX-2").te_prequant_repos == (
("fp8", "text_encoder", "unsloth/LTX-2-FP8"),
)
# The hosted filenames follow the repo naming convention the resolver derives.
assert te_prequant_repo_filename(
"unsloth/Qwen-Image-FP8", "text_encoder", "fp8"
) == "Qwen-Image-text_encoder-FP8.pt"
assert te_prequant_repo_filename(
"unsloth/FLUX.2-dev-FP8", "text_encoder", "fp8"
) == "FLUX.2-dev-text_encoder-FP8.pt"
assert te_prequant_repo_filename(
"unsloth/LTX-2-FP8", "text_encoder", "fp8"
) == "LTX-2-text_encoder-FP8.pt"
def test_cast_fp8_is_idempotent_on_precast_encoder():
"""A pre-cast encoder arrives with the layerwise hooks installed; the runtime re-apply in
quantize_text_encoders must be a no-op (re-registering the hook name raises, which made
the engaged cast report as failed and status show no TE quant)."""
import torch
from core.inference.diffusion_precision import _cast_fp8
target = types.SimpleNamespace(dtype = torch.bfloat16)
enc = torch.nn.Sequential(torch.nn.Linear(64, 64), torch.nn.LayerNorm(64))
_cast_fp8(enc, target)
assert enc[0].weight.dtype == torch.float8_e4m3fn
_cast_fp8(enc, target) # must not raise
assert enc[0].weight.dtype == torch.float8_e4m3fn
def test_builder_metadata_survives_weights_only_load(tmp_path):
"""The builder's checkpoint must load with torch.load(weights_only=True): version
metadata has to be plain str (a pickled TorchVersion object gets the whole artifact
rejected and the loader would silently fall back to the dense download)."""
import sys
import torch
scripts = Path(__file__).resolve().parents[3] / "scripts"
sys.path.insert(0, str(scripts))
try:
import build_te_prequant_checkpoint # noqa: F401 (import proves the module parses)
finally:
sys.path.remove(str(scripts))
ckpt = {
"format": TE_PREQUANT_FORMAT,
"metadata": {
"scheme": "fp8",
"component": "text_encoder",
"base_model_id": "Lightricks/LTX-2",
"te_class": "Gemma3ForConditionalGeneration",
"torch_version": str(torch.__version__),
"transformers_version": "0.0.0",
},
"state_dict": {"weight": torch.zeros(1)},
}
path = tmp_path / "te.pt"
torch.save(ckpt, path)
loaded = torch.load(path, weights_only = True, map_location = "cpu")
assert tpq._validate_checkpoint(loaded, "fp8", "text_encoder", "Lightricks/LTX-2", None)
# The regression: an unstringified TorchVersion in metadata must fail weights_only.
bad = dict(ckpt, metadata = dict(ckpt["metadata"], torch_version = torch.__version__))
bad_path = tmp_path / "bad.pt"
torch.save(bad, bad_path)
if not isinstance(torch.__version__, str):
with pytest.raises(Exception):
torch.load(bad_path, weights_only = True, map_location = "cpu")