Merge branch 'diffusion-auto-policy' into diffusion-fp16-accum
This commit is contained in:
commit
45e1c7eb33
4 changed files with 21 additions and 30 deletions
|
|
@ -1008,9 +1008,7 @@ class DiffusionBackend:
|
|||
cpu_offload,
|
||||
kind = kind,
|
||||
repo_id = repo_id,
|
||||
transformer_resident_override_mib = (
|
||||
candidate.transient_transformer_mib
|
||||
),
|
||||
transformer_resident_override_mib = (candidate.transient_transformer_mib),
|
||||
)
|
||||
if replanned.offload_policy == OFFLOAD_NONE:
|
||||
quant_plan = replanned
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@ import logging
|
|||
from dataclasses import dataclass
|
||||
from typing import Any, Optional
|
||||
|
||||
_MIB_PER_GB = 1000.0 ** 3 / (1024.0 * 1024.0) # component sizes below are decimal GB
|
||||
_MIB_PER_GB = 1000.0**3 / (1024.0 * 1024.0) # component sizes below are decimal GB
|
||||
|
||||
# Steady-state size of a torchao-quantised transformer relative to its bf16 weights:
|
||||
# int8 / fp8 store one byte per param plus per-row scales (~0.52x) with a little slack
|
||||
|
|
@ -161,7 +161,6 @@ def resolve_dense_quant_candidate(
|
|||
prequant_available = False
|
||||
try:
|
||||
from .diffusion_prequant import resolve_prequant_source
|
||||
|
||||
prequant_available = (
|
||||
resolve_prequant_source(fam, scheme, path_override = prequant_path) is not None
|
||||
)
|
||||
|
|
|
|||
|
|
@ -90,23 +90,25 @@ def test_estimate_unknown_family_or_scheme_returns_none():
|
|||
|
||||
|
||||
# ── candidate resolution (selector + prequant probe stubbed) ─────────────────
|
||||
def _patch_selector(monkeypatch, *, supported = True, scheme = "int8", prequant = None):
|
||||
def _patch_selector(
|
||||
monkeypatch,
|
||||
*,
|
||||
supported = True,
|
||||
scheme = "int8",
|
||||
prequant = None,
|
||||
):
|
||||
import core.inference.diffusion_transformer_quant as tq
|
||||
|
||||
monkeypatch.setattr(tq, "dense_transformer_supported", lambda target: supported)
|
||||
monkeypatch.setattr(tq, "select_transformer_quant_scheme", lambda target, req: scheme)
|
||||
import core.inference.diffusion_prequant as pq
|
||||
|
||||
monkeypatch.setattr(
|
||||
pq, "resolve_prequant_source", lambda fam, s, path_override = None: prequant
|
||||
)
|
||||
monkeypatch.setattr(pq, "resolve_prequant_source", lambda fam, s, path_override = None: prequant)
|
||||
|
||||
|
||||
def test_candidate_resolves_for_a_supported_request(monkeypatch):
|
||||
_patch_selector(monkeypatch, scheme = "int8")
|
||||
est = resolve_dense_quant_candidate(
|
||||
fam = _fam("z-image"), target = object(), requested = "auto"
|
||||
)
|
||||
est = resolve_dense_quant_candidate(fam = _fam("z-image"), target = object(), requested = "auto")
|
||||
assert isinstance(est, DenseQuantEstimate)
|
||||
assert est.scheme == "int8"
|
||||
assert est.transient_transformer_mib > est.steady_transformer_mib
|
||||
|
|
@ -115,44 +117,31 @@ def test_candidate_resolves_for_a_supported_request(monkeypatch):
|
|||
def test_candidate_none_when_request_is_off(monkeypatch):
|
||||
_patch_selector(monkeypatch)
|
||||
for off in (None, "", "none", "off"):
|
||||
assert (
|
||||
resolve_dense_quant_candidate(fam = _fam(), target = object(), requested = off)
|
||||
is None
|
||||
)
|
||||
assert resolve_dense_quant_candidate(fam = _fam(), target = object(), requested = off) is None
|
||||
|
||||
|
||||
def test_candidate_none_when_device_unsupported(monkeypatch):
|
||||
_patch_selector(monkeypatch, supported = False)
|
||||
assert (
|
||||
resolve_dense_quant_candidate(fam = _fam(), target = object(), requested = "auto")
|
||||
is None
|
||||
)
|
||||
assert resolve_dense_quant_candidate(fam = _fam(), target = object(), requested = "auto") is None
|
||||
|
||||
|
||||
def test_candidate_none_when_no_scheme_resolves(monkeypatch):
|
||||
_patch_selector(monkeypatch, scheme = None)
|
||||
assert (
|
||||
resolve_dense_quant_candidate(fam = _fam(), target = object(), requested = "auto")
|
||||
is None
|
||||
)
|
||||
assert resolve_dense_quant_candidate(fam = _fam(), target = object(), requested = "auto") is None
|
||||
|
||||
|
||||
def test_candidate_none_for_an_unlisted_family(monkeypatch):
|
||||
# No size entry -> no basis to re-plan; the loader keeps today's resident-only gate.
|
||||
_patch_selector(monkeypatch)
|
||||
assert (
|
||||
resolve_dense_quant_candidate(
|
||||
fam = _fam("not-a-family"), target = object(), requested = "auto"
|
||||
)
|
||||
resolve_dense_quant_candidate(fam = _fam("not-a-family"), target = object(), requested = "auto")
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_candidate_uses_prequant_transient_when_available(monkeypatch):
|
||||
_patch_selector(monkeypatch, prequant = object())
|
||||
est = resolve_dense_quant_candidate(
|
||||
fam = _fam("z-image"), target = object(), requested = "int8"
|
||||
)
|
||||
est = resolve_dense_quant_candidate(fam = _fam("z-image"), target = object(), requested = "int8")
|
||||
assert est is not None and est.prequant is True
|
||||
assert est.transient_transformer_mib == est.steady_transformer_mib
|
||||
|
||||
|
|
|
|||
|
|
@ -6,6 +6,8 @@ import json
|
|||
import sys
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from core.inference.diffusion_krea2 import (
|
||||
KREA2_FAMILY_NAME,
|
||||
_load_model_index,
|
||||
|
|
@ -187,6 +189,9 @@ def test_krea2_spec_registered_with_authors_targets():
|
|||
|
||||
|
||||
def test_krea2_collate_and_forward_roundtrip():
|
||||
# spec.forward imports Krea2Pipeline (prepare_position_ids), so this needs a real
|
||||
# diffusers install; CI hosts run the backend suite without one.
|
||||
pytest.importorskip("diffusers")
|
||||
import torch
|
||||
from core.training.diffusion_dit_trainer import _SPECS
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue