Gate mxfp8 DiT training precision before evicting resident GPU models
The start route's precision preflight folded bf16/int8/fp8 into the CUDA requirement but omitted mxfp8, so an mxfp8 request on a GPU-less host (or an older CUDA GPU without Blackwell) passed the preflight, evicted resident image and chat models, then raised only in the spawned trainer child. Mirror _resolve_base_precision: require CUDA for mxfp8 and re-check the Blackwell (sm100+) capability up front, so a doomed run is rejected before teardown.
This commit is contained in:
parent
eb80e66709
commit
22dad1df73
2 changed files with 42 additions and 10 deletions
|
|
@ -289,20 +289,21 @@ def bf16_unsupported_reason(resolved_family: str) -> Optional[str]:
|
|||
def training_precision_preflight_error(resolved_family: str, base_precision: str) -> Optional[str]:
|
||||
"""Reason the requested DiT precision cannot run on this host, else None -- checked by the
|
||||
start route BEFORE evicting resident GPU workloads (the trainer's own checks fire only in the
|
||||
child, after eviction). Three gates, all mirroring _resolve_base_precision so a doomed run is
|
||||
child, after eviction). Four gates, all mirroring _resolve_base_precision so a doomed run is
|
||||
rejected before teardown: the bf16-GPU requirement (bf16_unsupported_reason); the dense
|
||||
precisions (bf16/int8/fp8) requiring a CUDA GPU; and an explicit int8 needing a FUNCTIONAL
|
||||
torchao (its _int8_quantize_base has no fallback). Never raises."""
|
||||
precisions (bf16/int8/fp8/mxfp8) requiring a CUDA GPU; an explicit int8 needing a FUNCTIONAL
|
||||
torchao (its _int8_quantize_base has no fallback); and an explicit mxfp8 needing a Blackwell
|
||||
(sm100+) GPU (its MX GEMM has no kernel below sm100). Never raises."""
|
||||
reason = bf16_unsupported_reason(resolved_family)
|
||||
if reason:
|
||||
return reason
|
||||
fam = (resolved_family or "").strip().lower()
|
||||
mode = (base_precision or "").strip().lower()
|
||||
if fam in _DIT_TRAIN_FAMILIES and mode in ("bf16", "int8", "fp8"):
|
||||
if fam in _DIT_TRAIN_FAMILIES and mode in ("bf16", "int8", "fp8", "mxfp8"):
|
||||
# The DiT trainer's dense precisions all require CUDA (_resolve_base_precision rejects
|
||||
# bf16/int8/fp8 on device != "cuda"). bf16_unsupported_reason exempts a CPU-only host (the
|
||||
# fp32 fallback for import/unit tests), so without this a dense request on a GPU-less host
|
||||
# would pass the preflight, evict resident workloads, then raise only in the child.
|
||||
# bf16/int8/fp8/mxfp8 on device != "cuda"). bf16_unsupported_reason exempts a CPU-only host
|
||||
# (the fp32 fallback for import/unit tests), so without this a dense request on a GPU-less
|
||||
# host would pass the preflight, evict resident workloads, then raise only in the child.
|
||||
try:
|
||||
import torch
|
||||
has_cuda = torch.cuda.is_available()
|
||||
|
|
@ -318,6 +319,21 @@ def training_precision_preflight_error(resolved_family: str, base_precision: str
|
|||
"base_precision='int8' needs a functional torchao install; this host's torchao is "
|
||||
"missing or the non-functional Windows-ROCm stub. Use 'nf4', 'bf16', or 'auto'."
|
||||
)
|
||||
# mxfp8 needs Blackwell (sm100+): its MX GEMM has no kernel below sm100 and raises at the
|
||||
# first training step, AFTER a full dense-transformer load. Re-check here (mirroring
|
||||
# _resolve_base_precision) so a stale or direct client on an older CUDA GPU fails fast
|
||||
# before eviction instead of crashing mid-run.
|
||||
if mode == "mxfp8":
|
||||
try:
|
||||
import torch
|
||||
blackwell = torch.cuda.get_device_capability() >= (10, 0)
|
||||
except Exception: # noqa: BLE001 -- probe failure -> treat as unsupported, fail fast
|
||||
blackwell = False
|
||||
if not blackwell:
|
||||
return (
|
||||
"base_precision='mxfp8' needs a Blackwell (sm100+) GPU; this GPU is older. "
|
||||
"Use base_precision='bf16', 'int8', 'nf4', or 'auto'."
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -192,17 +192,33 @@ def test_training_precision_preflight_error(monkeypatch):
|
|||
assert training_precision_preflight_error("", "int8") is None
|
||||
|
||||
# On a CUDA-ABSENT host, bf16_unsupported_reason exempts CPU-only, but the DiT trainer's dense
|
||||
# precisions still require CUDA (mirroring _resolve_base_precision), so bf16/int8/fp8 for a DiT
|
||||
# family are rejected UP FRONT rather than after eviction. nf4/auto (and SDXL) still pass.
|
||||
# precisions still require CUDA (mirroring _resolve_base_precision), so bf16/int8/fp8/mxfp8 for
|
||||
# a DiT family are rejected UP FRONT rather than after eviction. nf4/auto (and SDXL) still pass.
|
||||
monkeypatch.setattr(common, "has_functional_torchao", lambda: True)
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: False)
|
||||
for dense in ("bf16", "int8", "fp8"):
|
||||
for dense in ("bf16", "int8", "fp8", "mxfp8"):
|
||||
reason = training_precision_preflight_error("flux.1", dense)
|
||||
assert reason is not None and "CUDA" in reason
|
||||
assert training_precision_preflight_error("flux.1", "nf4") is None
|
||||
assert training_precision_preflight_error("flux.1", "auto") is None
|
||||
assert training_precision_preflight_error("sdxl", "bf16") is None
|
||||
|
||||
# mxfp8 needs a Blackwell (sm100+) GPU: its MX GEMM has no kernel below sm100 and would raise at
|
||||
# the first training step, AFTER a full dense-transformer load. On a CUDA GPU that is older than
|
||||
# Blackwell the preflight rejects mxfp8 UP FRONT (mirroring _resolve_base_precision) so eviction
|
||||
# is skipped; other dense precisions on the same GPU still pass.
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
|
||||
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda *a, **k: (9, 0))
|
||||
reason = training_precision_preflight_error("flux.1", "mxfp8")
|
||||
assert reason is not None and "Blackwell" in reason
|
||||
assert training_precision_preflight_error("flux.1", "bf16") is None
|
||||
assert training_precision_preflight_error("flux.1", "fp8") is None
|
||||
|
||||
# On a Blackwell (sm100+) GPU mxfp8 is accepted, and it never gates a non-DiT (SDXL) family.
|
||||
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda *a, **k: (10, 0))
|
||||
assert training_precision_preflight_error("flux.1", "mxfp8") is None
|
||||
assert training_precision_preflight_error("sdxl", "mxfp8") is None
|
||||
|
||||
|
||||
def test_family_train_infos_empties_dit_modes_on_non_bf16(monkeypatch):
|
||||
# On a non-bf16 GPU the start route rejects EVERY DiT family (even nf4), so /info must not
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue