mirror of
https://github.com/unslothai/unsloth.git
synced 2026-08-25 08:42:25 +00:00
* Fix FLUX.2 Klein LoRA training quality * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Fix FLUX.2 Klein training review gaps * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Fix diffusion training mirror preflight * Drop base_specs on a DiT block so the hardware reason still reaches the UI A dit_block clears the family-level chip fields (params, qlora_vram_gb, gated, note) precisely so FamilyFacts falls through to vram_note, which carries the actionable reason: no CUDA, or no native bf16 on this GPU. base_specs was still published though, and the per-base entry wins in resolveDiffusionTrainingFacts, so selecting Klein base-9B on a blocked host put the 9B and 18 GB chips straight back. FamilyFacts renders vram_note only when there are no chips, so the user lost the explanation and got a generic refusal plus a size they cannot act on. Clear base_specs on the same condition. The overlay only ever feeds those chips (resolveDiffusionTrainingFacts is its sole consumer), so nothing functional depends on it being present. Covered by a test that pins a bf16-unsupported host and asserts both halves: the overlay is empty and vram_note still carries the reason. It also asserts a per-base overlay exists on an unblocked host first, so it cannot pass vacuously. Verified it fails with the fix reverted. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Keep the mirror out of the gate and out of the resume identity Two consequences of routing the fetch through prefer_ungated_mirror that the mirror was not meant to have. The gated-repo name check ran on the canonical id while the start route preflighted the fetch repo. For FLUX.1-dev or FLUX.2-dev without a token that means the route probes the public mirror, answers 200, frees the resident models, and only then does the child refuse the run by name. It now checks the repo the run will actually fetch, so the two agree: a mirror-backed run proceeds, and a genuinely gated fetch still fails in the route, before anything is evicted. base_revision was reading the fetch repo too. Which repo that is depends on local cache state, so the same base records a different rev- value once the upstream snapshot is evicted or the run moves to another machine, and mismatch_reason then refuses the checkpoint as a different base revision though the weights are byte identical. Every checkpoint written before mirrors existed holds the canonical value, so those would be refused as well. Both identity_for_config and the post-load re-read go back to the canonical id; a canonical repo with no local ref records "unresolved", which is already treated as not comparable. Tests, both mutation checked: the gate assertion fails when the canonical id is restored, and the revision assertion fails when the fetch repo is. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Pair the recorded base revision with the repo it was read from A mirror-backed run resolves the revision from the mirror, so recording the canonical repo's revision instead lost validation entirely: source_revision(canonical) is "unresolved" on exactly the hosts a mirror is selected for, and mismatch_reason skips non-comparable revisions, so the check silently never fired. Record both halves and compare them only when the repos match, which still catches a moved mirror and no longer refuses a legitimate cross-repo resume. Bundles written before the field exist keep comparing as they did. Also take the mirror on a token-less run even when the vendor repo is cached: only gated repos are in the mirror table, so a mirror existing means the upstream needs credentials this run does not have. UNSLOTH_DIFFUSION_NO_MIRROR still pins the vendor repo. And preselect the training base a family pairs with a loaded distilled checkpoint: the distilled half is never in base_repos, so opening Train with the 9B model loaded seeded the 4B base. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Only override a cached base when its upstream is actually gated The mirror table stopped being gated-only when the redistributable bases joined it, so "a mirror exists" no longer implies "the upstream needs credentials". Klein base-4B is both a default trainable base and mirrored, and the Hub serves it anonymously, so a token-less run with a complete local cache was re-pulling it from the mirror and an offline run failed outright. Split the table rather than keep a second list in sync: the 12 genuinely gated pairs stay in _GATED_MIRROR_PAIRS, the redistributable ones move to _UNGATED_MIRROR_PAIRS, and mirror_repo and canonical_base build from the union so redirection and the base-keyed table normalisation are unchanged. upstream_is_gated reads the gated half, and the token-less override now asks it. * Keep a local clone named like a gated base out of the mirror override * Carry the local-clone and mirror exceptions through every base gate --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Daniel Han <danielhanchen@gmail.com>
492 lines
22 KiB
Python
492 lines
22 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 / FLUX.2 / Qwen-Image / Z-Image / LTX-2).
|
|
|
|
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 types
|
|
|
|
import pytest
|
|
|
|
from core.training.diffusion_dit_trainer import (
|
|
_FLUX2_DEV_TARGETS,
|
|
_FLUX2_KLEIN_TARGETS,
|
|
_FLUX_TARGETS,
|
|
_GATED_TRAIN_REPOS,
|
|
_QWEN_TARGETS,
|
|
_SPECS,
|
|
_ZIMAGE_TARGETS,
|
|
_apply_mxfp8_training,
|
|
_assert_gated_access,
|
|
_mx_module_filter,
|
|
_repo_is_prequantized,
|
|
_resolve_base_precision,
|
|
_select_lora_targets,
|
|
_should_compile,
|
|
run_dit_lora_training,
|
|
)
|
|
from core.training.diffusion_train_common import (
|
|
DEFAULT_LORA_TARGETS,
|
|
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.2-klein",
|
|
"flux.2-dev",
|
|
# The first VIDEO family; its own assertions live in test_diffusion_dit_trainer_ltx2.
|
|
"ltx-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, Krea 2 and both FLUX.2 variants 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
|
|
assert _SPECS["flux.2-klein"].force_bf16 is True
|
|
assert _SPECS["flux.2-dev"].force_bf16 is True
|
|
|
|
|
|
def test_flux2_specs_share_targets_and_split_conditioners():
|
|
# dev and Klein share the transformer class but have different single-block counts.
|
|
klein, dev = _SPECS["flux.2-klein"], _SPECS["flux.2-dev"]
|
|
assert klein.lora_targets == _FLUX2_KLEIN_TARGETS
|
|
assert dev.lora_targets == _FLUX2_DEV_TARGETS
|
|
# The upstream trainers pair the fused input with every plain single-stream output projection.
|
|
assert "to_qkv_mlp_proj" in _FLUX2_KLEIN_TARGETS
|
|
assert "to_out.0" in _FLUX2_KLEIN_TARGETS
|
|
assert "single_transformer_blocks.23.attn.to_out" in _FLUX2_KLEIN_TARGETS
|
|
assert "single_transformer_blocks.24.attn.to_out" not in _FLUX2_KLEIN_TARGETS
|
|
assert "single_transformer_blocks.47.attn.to_out" in _FLUX2_DEV_TARGETS
|
|
assert klein.load_conditioners is not dev.load_conditioners
|
|
assert klein.save is not dev.save
|
|
assert klein.load_transformer is dev.load_transformer
|
|
# The Mistral stack makes dev far heavier than the 4B Klein.
|
|
assert dev.dense_bf16_gb > klein.dense_bf16_gb
|
|
|
|
|
|
def test_select_lora_targets_uses_family_default_for_generic_config():
|
|
# normalized() fills lora_target_modules with DEFAULT_LORA_TARGETS, so that value must resolve to the family's targets, not stay on the SDXL list.
|
|
assert _select_lora_targets(DEFAULT_LORA_TARGETS, _FLUX_TARGETS) == _FLUX_TARGETS
|
|
assert _select_lora_targets(DEFAULT_LORA_TARGETS, _QWEN_TARGETS) == _QWEN_TARGETS
|
|
assert _select_lora_targets(DEFAULT_LORA_TARGETS, _ZIMAGE_TARGETS) == _ZIMAGE_TARGETS
|
|
|
|
|
|
def test_select_lora_targets_explicit_override_wins():
|
|
# Any OTHER explicit tuple is a deliberate override and must win over the family spec.
|
|
override = ("to_q", "to_k")
|
|
assert _select_lora_targets(override, _FLUX_TARGETS) == override
|
|
# The default request path (config carrying the generic default) reaches the spec.
|
|
cfg = DiffusionLoraConfig(
|
|
base_model = "black-forest-labs/FLUX.1-dev", data_dir = "d", output_dir = "o"
|
|
).normalized()
|
|
assert cfg.lora_target_modules == DEFAULT_LORA_TARGETS
|
|
assert (
|
|
_select_lora_targets(cfg.lora_target_modules, _SPECS["flux.1"].lora_targets)
|
|
== _FLUX_TARGETS
|
|
)
|
|
|
|
|
|
@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_flux2_rejects_fp16_before_loading():
|
|
# Both FLUX.2 variants resolve from their repo names and are bf16-only, so an explicit fp16 fails in normalized(). Klein's base is ungated, exercising the guard directly.
|
|
ok = DiffusionLoraConfig(
|
|
base_model = "black-forest-labs/FLUX.2-klein-4B", data_dir = "d", output_dir = "o"
|
|
).normalized()
|
|
assert ok.resolved_family == "flux.2-klein"
|
|
assert (
|
|
DiffusionLoraConfig(base_model = "black-forest-labs/FLUX.2-dev", data_dir = "d", output_dir = "o")
|
|
.normalized()
|
|
.resolved_family
|
|
== "flux.2-dev"
|
|
)
|
|
cfg = DiffusionLoraConfig(
|
|
base_model = "black-forest-labs/FLUX.2-klein-4B",
|
|
data_dir = "does-not-exist",
|
|
output_dir = "o",
|
|
mixed_precision = "fp16",
|
|
)
|
|
with pytest.raises(ValueError, match = "bf16"):
|
|
run_dit_lora_training(cfg)
|
|
|
|
|
|
def test_flux2_bases_pass_the_trusted_base_gate():
|
|
# The FLUX.2 bases are training-side additions to the loader's trust allowlist, so the pre-download trust gate must accept them.
|
|
from core.training.diffusion_train_common import _assert_trusted_base_model
|
|
|
|
_assert_trusted_base_model("black-forest-labs/FLUX.2-klein-base-4B")
|
|
_assert_trusted_base_model("black-forest-labs/FLUX.2-klein-base-9B")
|
|
_assert_trusted_base_model("black-forest-labs/FLUX.2-klein-4B")
|
|
_assert_trusted_base_model("black-forest-labs/FLUX.2-dev")
|
|
with pytest.raises(ValueError, match = "untrusted"):
|
|
_assert_trusted_base_model("someone/random-flux2-finetune")
|
|
|
|
|
|
def test_zimage_offers_the_undistilled_base_the_upstream_recipe_trains_on():
|
|
# examples/dreambooth/README_z_image.md trains on Tongyi-MAI/Z-Image, not the distilled Turbo,
|
|
# and the trust gate refused that id until it joined the allowlist. The nf4 Turbo stays first
|
|
# so it remains the picker's default.
|
|
from core.inference.diffusion_families import detect_family
|
|
from core.training.diffusion_train_common import _assert_trusted_base_model
|
|
|
|
fam = detect_family("Tongyi-MAI/Z-Image")
|
|
assert fam is not None and fam.name == "z-image"
|
|
assert fam.train_base_repos == (
|
|
"unsloth/Z-Image-Turbo-unsloth-bnb-4bit",
|
|
"Tongyi-MAI/Z-Image-Turbo",
|
|
"Tongyi-MAI/Z-Image",
|
|
)
|
|
_assert_trusted_base_model("Tongyi-MAI/Z-Image")
|
|
with pytest.raises(ValueError, match = "untrusted"):
|
|
_assert_trusted_base_model("someone/random-z-image-finetune")
|
|
# The upstream script's target list; the family spec must already match it.
|
|
assert _SPECS["z-image"].lora_targets == ("to_q", "to_k", "to_v", "to_out.0")
|
|
# No deploy pairing: an adapter previews on whichever checkpoint it trained on. A family-wide
|
|
# one would also rewrite the nf4 Turbo base, sending a QLoRA run's preview to a dense fp32 load.
|
|
assert fam.deploy_base_repo is None
|
|
|
|
|
|
def test_every_train_base_is_deployable_as_an_inference_pipeline():
|
|
# "Deploy to Create" reloads the trained-on base (or the family's deploy_base) through /images/load as a PIPELINE, gated on
|
|
# _is_trusted_diffusion_repo, so an advertised training base failing that gate makes Deploy 400 for every adapter.
|
|
from core.inference.diffusion import _is_trusted_diffusion_repo
|
|
from core.inference.diffusion_families import _FAMILIES
|
|
for fam in _FAMILIES:
|
|
if not fam.trainable:
|
|
continue
|
|
for base in fam.train_base_repos:
|
|
deploy_base = fam.deploy_base_for(base)
|
|
assert _is_trusted_diffusion_repo(
|
|
deploy_base
|
|
), f"{fam.name}: deploy base {deploy_base!r} is not loadable for inference"
|
|
|
|
|
|
def test_gated_access_requires_token():
|
|
assert "black-forest-labs/flux.1-dev" in _GATED_TRAIN_REPOS
|
|
assert "black-forest-labs/flux.2-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 pytest.raises(ValueError, match = "gated"):
|
|
_assert_gated_access("black-forest-labs/FLUX.2-dev", None)
|
|
# 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("black-forest-labs/FLUX.2-dev", "hf_realtoken")
|
|
_assert_gated_access("Tongyi-MAI/Z-Image-Turbo", None)
|
|
_assert_gated_access("black-forest-labs/FLUX.2-klein-4B", None) # Klein is open
|
|
|
|
|
|
def test_the_gate_lets_a_local_clone_named_like_a_gated_repo_through(monkeypatch, tmp_path):
|
|
"""A directory on disk carries no gate, whatever it is called.
|
|
|
|
A base can be a relative clone named exactly like the vendor repo, which the loaders and the
|
|
token-less mirror override both resolve on disk. Matching \`_GATED_TRAIN_REPOS\` by name alone
|
|
refused that layout without a token, for weights the run never fetches.
|
|
"""
|
|
local = "black-forest-labs/FLUX.1-dev"
|
|
assert local.lower() in _GATED_TRAIN_REPOS, "precondition: the name is gated"
|
|
monkeypatch.chdir(tmp_path)
|
|
(tmp_path / local).mkdir(parents = True)
|
|
|
|
_assert_gated_access(local, None)
|
|
|
|
|
|
def test_the_gate_reads_the_repo_the_run_will_fetch(monkeypatch, tmp_path):
|
|
"""A gated base redirected to its ungated mirror must not be refused by name.
|
|
|
|
The start route preflights the FETCH repo, so a child that checked the canonical id
|
|
would raise for a request the route had already answered 200 to, after freeing the
|
|
resident models: a dead job instead of a fast 400.
|
|
"""
|
|
from core.inference import diffusion_families
|
|
|
|
seen: list[str] = []
|
|
monkeypatch.setattr(
|
|
"core.training.diffusion_dit_trainer._assert_gated_access",
|
|
lambda base, token: seen.append(base),
|
|
)
|
|
monkeypatch.setattr(
|
|
diffusion_families,
|
|
"prefer_ungated_mirror",
|
|
lambda base, token = None: "unsloth/FLUX.1-dev"
|
|
if base.lower() == "black-forest-labs/flux.1-dev"
|
|
else base,
|
|
)
|
|
cfg = DiffusionLoraConfig(
|
|
base_model = "black-forest-labs/FLUX.1-dev",
|
|
data_dir = str(tmp_path / "empty"), # the next step after the gate, so it stops here
|
|
output_dir = str(tmp_path / "out"),
|
|
)
|
|
with pytest.raises(Exception): # noqa: B017, PT011 -- the dataset, not the gate
|
|
run_dit_lora_training(cfg)
|
|
assert seen == ["unsloth/FLUX.1-dev"]
|
|
|
|
# Control: with no mirror at all the canonical id is still what gets checked, so a
|
|
# genuinely gated fetch without a token keeps failing here rather than mid-download.
|
|
# mirror_repo has to go too: a token-less run overrides the cache preference on any repo
|
|
# the mirror table covers, so stubbing only the preference would still redirect.
|
|
seen.clear()
|
|
monkeypatch.setattr(diffusion_families, "prefer_ungated_mirror", lambda base, token = None: base)
|
|
monkeypatch.setattr(diffusion_families, "mirror_repo", lambda base: None)
|
|
with pytest.raises(Exception): # noqa: B017, PT011
|
|
run_dit_lora_training(cfg)
|
|
assert seen == ["black-forest-labs/FLUX.1-dev"]
|
|
|
|
|
|
def test_family_train_infos_lists_dit_families(dit_train_host):
|
|
infos = {i["name"]: i for i in family_train_infos()}
|
|
for fam in ("sdxl", "flux.1", "qwen-image", "z-image", "flux.2-klein", "flux.2-dev"):
|
|
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 bases are the gated dev repos; their notes flag the license requirement.
|
|
assert infos["flux.1"]["default_base"] == "black-forest-labs/FLUX.1-dev"
|
|
assert "gated" in infos["flux.1"]["vram_note"].lower()
|
|
assert infos["flux.2-dev"]["default_base"] == "black-forest-labs/FLUX.2-dev"
|
|
assert "gated" in infos["flux.2-dev"]["vram_note"].lower()
|
|
# Klein trains on the undistilled bases and deploys each size on its distilled partner.
|
|
klein = infos["flux.2-klein"]
|
|
assert klein["default_base"] == "black-forest-labs/FLUX.2-klein-base-4B"
|
|
assert klein["base_repos"] == [
|
|
"black-forest-labs/FLUX.2-klein-base-4B",
|
|
"black-forest-labs/FLUX.2-klein-base-9B",
|
|
]
|
|
assert klein["deploy_bases"]["black-forest-labs/FLUX.2-klein-base-4B"] == (
|
|
"black-forest-labs/FLUX.2-klein-4B"
|
|
)
|
|
assert klein["deploy_bases"]["black-forest-labs/FLUX.2-klein-base-9B"] == (
|
|
"black-forest-labs/FLUX.2-klein-9B"
|
|
)
|
|
assert klein["deploy_bases"]["unsloth/FLUX.2-klein-base-9B"] == ("unsloth/FLUX.2-klein-9B")
|
|
assert klein["base_specs"]["black-forest-labs/FLUX.2-klein-base-9B"] == {
|
|
"params": "9B",
|
|
"qlora_vram_gb": 18,
|
|
}
|
|
assert klein["base_specs"]["unsloth/FLUX.2-klein-base-9B"] == {
|
|
"params": "9B",
|
|
"qlora_vram_gb": 18,
|
|
}
|
|
assert "black-forest-labs/FLUX.2-klein-base-4B" not in klein["base_specs"]
|
|
assert "gated" not in infos["flux.2-klein"]["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, dit_train_host
|
|
):
|
|
# Regional compile applies to every family (SDXL compiles its U-Net blocks too) but base_precision stays DiT-only, so SDXL
|
|
# advertises no precision modes while z-image keeps its own. Pin the list so the assertion holds on any host GPU.
|
|
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,
|
|
bias = False,
|
|
):
|
|
import torch.nn as nn
|
|
return nn.Linear(in_features, out_features, bias = bias)
|
|
|
|
|
|
def test_mx_module_filter_accepts_dense_block_linear():
|
|
# A bias-free 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_biased_linear():
|
|
# The torchao 0.17 MX training path drops the bias, so an mxfp8'd biased FROZEN linear would corrupt the base output the LoRA regresses against.
|
|
assert _mx_module_filter(_linear(3072, 3072, bias = True), "blocks.0.ff.up") is False
|
|
|
|
|
|
def test_resolve_base_precision_explicit_mxfp8_requires_blackwell(monkeypatch):
|
|
# An explicit mxfp8 on a non-Blackwell CUDA GPU must fail fast: the MX GEMM has no kernel below sm100 and would crash at the first step, after a full dense load.
|
|
import torch
|
|
|
|
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda *a, **k: (8, 9))
|
|
cfg = types.SimpleNamespace(base_precision = "mxfp8", mixed_precision = "bf16", base_model = "x")
|
|
with pytest.raises(ValueError, match = "Blackwell"):
|
|
_resolve_base_precision(cfg, None, "cuda")
|
|
|
|
|
|
def test_resolve_base_precision_explicit_mxfp8_ok_on_blackwell(monkeypatch):
|
|
import torch
|
|
|
|
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda *a, **k: (10, 0))
|
|
cfg = types.SimpleNamespace(base_precision = "mxfp8", mixed_precision = "bf16", base_model = "x")
|
|
assert _resolve_base_precision(cfg, None, "cuda") == "mxfp8"
|
|
|
|
|
|
def test_mx_module_filter_skips_lora_and_proj_out():
|
|
# LoRA-owned modules and the output projection are excluded, mirroring the fp8 filter.
|
|
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 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 both API revisions' imports to raise and assert one warning naming mxfp8.
|
|
monkeypatch.setitem(sys.modules, "torchao.prototype.mx_formats", None)
|
|
monkeypatch.setitem(sys.modules, "torchao.prototype.moe_training.config", 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 test_mxfp8_training_config_falls_back_to_the_torchao_0_17_api(monkeypatch):
|
|
# torchao 0.17 replaced prototype.mx_formats.MXLinearConfig with the MXFP8TrainingOpConfig recipe API, so the config helper must fall back or mxfp8 silently trains dense bf16.
|
|
from types import SimpleNamespace
|
|
|
|
from core.training.diffusion_dit_trainer import _mxfp8_training_config
|
|
|
|
calls = {}
|
|
|
|
class _Recipe:
|
|
MXFP8_RCEIL = "mxfp8_rceil"
|
|
|
|
class _OpConfig:
|
|
@staticmethod
|
|
def from_recipe(recipe):
|
|
calls["recipe"] = recipe
|
|
return "cfg-0.17"
|
|
|
|
fake_config = SimpleNamespace(MXFP8TrainingOpConfig = _OpConfig, MXFP8TrainingRecipe = _Recipe)
|
|
monkeypatch.setitem(sys.modules, "torchao.prototype.mx_formats", None)
|
|
monkeypatch.setitem(
|
|
sys.modules, "torchao.prototype.moe_training", SimpleNamespace(config = fake_config)
|
|
)
|
|
monkeypatch.setitem(sys.modules, "torchao.prototype.moe_training.config", fake_config)
|
|
assert _mxfp8_training_config() == "cfg-0.17"
|
|
assert calls["recipe"] == _Recipe.MXFP8_RCEIL
|
|
|
|
|
|
def _patch_capability(monkeypatch, capability):
|
|
# Drive train_precision_modes' GPU probe: pretend CUDA is present at the given capability (fp8 needs sm89+, mxfp8 sm100+).
|
|
# torchao is stubbed functional so these test the CAPABILITY gate, and is_bf16_supported is stubbed True (Ada/Blackwell always are).
|
|
import torch
|
|
|
|
import core.training.diffusion_train_common as dtc
|
|
|
|
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
|
|
monkeypatch.setattr(torch.cuda, "is_bf16_supported", lambda *a, **k: True)
|
|
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda *a, **k: capability)
|
|
monkeypatch.setattr(dtc, "has_functional_torchao", lambda: True)
|
|
|
|
|
|
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
|
|
|
|
|
|
def test_train_precision_modes_pre_ampere_is_nf4_only(monkeypatch):
|
|
# A pre-Ampere GPU EMULATES bf16 with no native tensor cores and the DiT trainer requires native bf16, so /info must offer
|
|
# nf4 only, else it advertises a start that evicts resident models and then fails the trainer's bf16 guard.
|
|
import torch
|
|
|
|
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
|
|
monkeypatch.setattr(
|
|
torch.cuda, "is_bf16_supported", lambda *a, **k: True
|
|
) # emulation reports True
|
|
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda *a, **k: (7, 5)) # Turing
|
|
modes, recommended = train_precision_modes()
|
|
assert modes == ["nf4"]
|
|
assert recommended == "nf4"
|