Merge branch 'diffusion-auto-policy' into diffusion-fp16-accum

This commit is contained in:
Daniel Han 2026-07-04 07:27:07 +00:00
commit 45e1c7eb33
4 changed files with 21 additions and 30 deletions

View file

@ -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

View file

@ -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
)

View file

@ -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

View file

@ -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