unsloth/studio/backend/tests/test_diffusion_dit_trainer.py
Daniel Han cb9247e537 Add mxfp8 training base precision and SDXL U-Net regional compile
- base_precision="mxfp8": torchao MX block-scaled float8 compute on the frozen
  base linears (Blackwell sm100+, cuBLAS kernels). Applied after add_adapter like
  fp8, never fatal, weights stay bf16 in memory. Measured 1.16x over compiled
  bf16 on Z-Image at 1024px batch 4 (16k tokens/step); a wash at small token
  counts, so it stays an explicit opt-in and auto never picks it.
- SDXL: regionally compile the U-Net's BasicTransformerBlocks through the same
  never-fatal wrapper the DiT trainer uses. 1.35x steady state at 1024px batch 4
  with same-seed loss parity (~1e-5 per step) and unchanged peak VRAM; ~30 s
  one-time warmup. Steady-state samples/sec now excludes step 1, matching the
  DiT trainer.
- /info: mxfp8 advertised only on sm100+; supports_compile now true for sdxl.
- NVFP4 training: not available in torchao 0.16 (no autograd path, no training
  recipe), so NVFP4 stays an inference-only quant for now.

193 diffusion backend tests green; frontend build clean.
2026-07-03 19:13:03 +00:00

202 lines
8.2 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
"""Unit tests for the flow-matching DiT LoRA trainer (FLUX.1 / Qwen-Image / Z-Image).
CPU-only: cover family resolution, the per-family spec table, the QLoRA prequant
heuristic, the bf16-only guard, and the gated-repo name check. The full training loop is
exercised by the live GPU smokes, not here."""
from __future__ import annotations
import sys
import pytest
from core.training.diffusion_dit_trainer import (
_GATED_TRAIN_REPOS,
_SPECS,
_apply_mxfp8_training,
_assert_gated_access,
_mx_module_filter,
_repo_is_prequantized,
_should_compile,
run_dit_lora_training,
)
from core.training.diffusion_train_common import (
DiffusionLoraConfig,
family_train_infos,
train_precision_modes,
)
def test_specs_cover_the_dit_families():
assert set(_SPECS) == {"flux.1", "qwen-image", "z-image", "krea-2"}
# FLUX / Qwen share the added-kv attention target set; Z-Image and Krea 2 are
# single-stream.
assert "add_q_proj" in _SPECS["flux.1"].lora_targets
assert "add_q_proj" in _SPECS["qwen-image"].lora_targets
assert "add_q_proj" not in _SPECS["z-image"].lora_targets
assert "add_q_proj" not in _SPECS["krea-2"].lora_targets
# Z-Image, Qwen and Krea 2 are bf16-only.
assert _SPECS["z-image"].force_bf16 is True
assert _SPECS["qwen-image"].force_bf16 is True
assert _SPECS["krea-2"].force_bf16 is True
@pytest.mark.parametrize(
"repo, expected",
[
("unsloth/Qwen-Image-2512-unsloth-bnb-4bit", True),
("unsloth/Z-Image-Turbo-unsloth-bnb-4bit", True),
("some/model-int4", True),
("black-forest-labs/FLUX.1-dev", False),
("Tongyi-MAI/Z-Image-Turbo", False),
],
)
def test_prequant_heuristic(repo, expected):
assert _repo_is_prequantized(repo) is expected
def test_zimage_rejects_fp16_before_loading():
# bf16-only families must refuse an explicit fp16 request up front (no model load).
cfg = DiffusionLoraConfig(
base_model = "Tongyi-MAI/Z-Image-Turbo",
data_dir = "does-not-exist",
output_dir = "o",
mixed_precision = "fp16",
)
with pytest.raises(ValueError, match = "bf16"):
run_dit_lora_training(cfg)
def test_gated_access_requires_token():
assert "black-forest-labs/flux.1-dev" in _GATED_TRAIN_REPOS
# No token -> clear, actionable error before any download.
with pytest.raises(ValueError, match = "gated"):
_assert_gated_access("black-forest-labs/FLUX.1-dev", None)
with pytest.raises(ValueError, match = "gated"):
_assert_gated_access("black-forest-labs/FLUX.1-dev", " ")
# With a token, or for a non-gated repo, it is a no-op.
_assert_gated_access("black-forest-labs/FLUX.1-dev", "hf_realtoken")
_assert_gated_access("Tongyi-MAI/Z-Image-Turbo", None)
def test_family_train_infos_lists_dit_families():
infos = {i["name"]: i for i in family_train_infos()}
for fam in ("sdxl", "flux.1", "qwen-image", "z-image"):
assert fam in infos, f"{fam} missing from family_train_infos"
assert infos[fam]["default_base"]
assert infos[fam]["base_repos"]
assert "resolution" in infos[fam]["defaults"]
# FLUX default base is the gated dev repo; its note flags the license requirement.
assert infos["flux.1"]["default_base"] == "black-forest-labs/FLUX.1-dev"
assert "gated" in infos["flux.1"]["vram_note"].lower()
# Z-Image defaults to the prequant nf4 repo for QLoRA.
assert "4bit" in infos["z-image"]["default_base"].lower()
def test_family_train_infos_sdxl_supports_compile_without_precision_modes(monkeypatch):
# Regional compile now applies to every family (the SDXL trainer compiles its U-Net
# blocks too), but base_precision stays DiT-only, so SDXL advertises no precision modes
# while a DiT family (z-image) keeps its own. Pin the precision list so the assertion
# holds regardless of the test host's GPU capability.
import core.training.diffusion_train_common as dtc
monkeypatch.setattr(dtc, "train_precision_modes", lambda: (["nf4", "bf16", "auto"], "auto"))
infos = {i["name"]: i for i in family_train_infos()}
assert infos["sdxl"]["supports_compile"] is True
assert infos["sdxl"]["precision_modes"] == []
assert infos["z-image"]["supports_compile"] is True
assert infos["z-image"]["precision_modes"] == ["nf4", "bf16", "auto"]
# ── mxfp8 base precision (DiT dense speed mode) ───────────────────────────────
def _linear(in_features, out_features):
import torch.nn as nn
return nn.Linear(in_features, out_features)
def test_mx_module_filter_accepts_dense_block_linear():
# A 3072x3072 attention/FFN linear at a normal block fqn is a valid mxfp8 target.
assert _mx_module_filter(_linear(3072, 3072), "blocks.0.ff.up") is True
def test_mx_module_filter_skips_lora_and_proj_out():
# LoRA-owned modules (adapters stay high precision) and the output projection are
# excluded, mirroring the fp8 filter's guards.
lin = _linear(3072, 3072)
assert _mx_module_filter(lin, "blocks.0.attn.to_q.lora_A.default") is False
assert _mx_module_filter(lin, "proj_out") is False
assert _mx_module_filter(lin, "x.proj_out.y") is False
def test_mx_module_filter_rejects_non_block_aligned_dims():
# MX block scaling tiles 32-wide, so a dim not divisible by 32 (3000) is rejected.
assert _mx_module_filter(_linear(3000, 3072), "blocks.0.ff.up") is False
def test_mx_module_filter_rejects_non_linear():
import torch.nn as nn
# A non-Linear module is never a target even if it exposes matching feature counts.
assert _mx_module_filter(nn.LayerNorm(3072), "blocks.0.norm") is False
def test_should_compile_auto_mxfp8_on_cuda():
# auto compiles the dense speed modes on cuda; int8 stays eager (torchao subclass);
# an explicit "off" wins over the mode.
cfg = DiffusionLoraConfig(base_model = "b", data_dir = "d", output_dir = "o")
assert _should_compile(cfg, False, "cuda", base_precision = "mxfp8") is True
assert _should_compile(cfg, False, "cuda", base_precision = "int8") is False
off = DiffusionLoraConfig(
base_model = "b", data_dir = "d", output_dir = "o", compile_transformer = "off"
)
assert _should_compile(off, False, "cuda", base_precision = "mxfp8") is False
def test_apply_mxfp8_training_failure_falls_back_with_warning(monkeypatch):
# An unavailable torchao MX path must never be fatal: force the import to raise, then
# assert the helper returns False and emits exactly one warning naming mxfp8.
monkeypatch.setitem(sys.modules, "torchao.prototype.mx_formats", None)
events = []
ok = _apply_mxfp8_training(object(), lambda e: events.append(e))
assert ok is False
warnings = [e for e in events if e["type"] == "warning"]
assert len(warnings) == 1
assert "mxfp8" in warnings[0]["message"]
def _patch_capability(monkeypatch, capability):
# Drive train_precision_modes' GPU probe: pretend CUDA is present at the given tensor
# core capability (fp8 needs sm89+, mxfp8 needs sm100+).
import torch
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda *a, **k: capability)
def test_train_precision_modes_blackwell_lists_mxfp8(monkeypatch):
# sm100 (Blackwell) exposes both fp8 and mxfp8, ordered before the "auto" pick.
_patch_capability(monkeypatch, (10, 0))
modes, recommended = train_precision_modes()
assert "mxfp8" in modes and "fp8" in modes
assert modes.index("mxfp8") < modes.index("auto")
assert modes.index("fp8") < modes.index("auto")
assert recommended == "auto"
def test_train_precision_modes_ada_has_fp8_without_mxfp8(monkeypatch):
# sm89 (Ada) is fp8-capable but not block-scaled mxfp8-capable.
_patch_capability(monkeypatch, (8, 9))
modes, _ = train_precision_modes()
assert "fp8" in modes
assert "mxfp8" not in modes
def test_train_precision_modes_newer_blackwell_has_mxfp8(monkeypatch):
# Any capability >= sm100 keeps mxfp8 (sm120 here).
_patch_capability(monkeypatch, (12, 0))
modes, _ = train_precision_modes()
assert "mxfp8" in modes