Stability improvements to model offloading. Added D-OPSD bleed loss as well.
This commit is contained in:
@@ -1153,6 +1153,9 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
|
||||
if self.dopsd:
|
||||
self.dopsd_self_ref = True
|
||||
self.require_pixel_tensor_cache = True
|
||||
self.dopsd_bleed_strength = float(
|
||||
self.model_config.model_kwargs.get("dopsd_bleed_strength", 1.0)
|
||||
)
|
||||
|
||||
def _dit_component(self) -> str:
|
||||
partition = str(
|
||||
|
||||
@@ -188,6 +188,10 @@ class ZImageModel(BaseModel):
|
||||
use_comfy_weights=self.model_config.model_kwargs.get(
|
||||
"use_comfy_weights", True
|
||||
),
|
||||
# comfy candidate ranking hint only (aitk_post_load quantizes):
|
||||
# a matching pre-quantized file loads with no requant work
|
||||
qtype=self.model_config.qtype if self.model_config.quantize else None,
|
||||
quantize_on_load=False,
|
||||
)
|
||||
flush()
|
||||
|
||||
|
||||
@@ -556,6 +556,7 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
noise_pred = noise_pred * self.train_config.pred_scaler
|
||||
|
||||
target = None
|
||||
dopsd_normal_target = None
|
||||
|
||||
if self.train_config.target_noise_multiplier != 1.0:
|
||||
noise = noise * self.train_config.target_noise_multiplier
|
||||
@@ -623,6 +624,18 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
assert not self.train_config.train_turbo
|
||||
# matching adapter prediction
|
||||
target = prior_pred
|
||||
if getattr(self.sd, 'dopsd_self_ref', False):
|
||||
# D-OPSD bleed: also train against the normal (non-teacher) target
|
||||
if hasattr(self.sd, 'get_loss_target'):
|
||||
dopsd_normal_target = self.sd.get_loss_target(
|
||||
noise=noise,
|
||||
batch=batch,
|
||||
timesteps=timesteps,
|
||||
).detach()
|
||||
elif self.sd.is_flow_matching:
|
||||
dopsd_normal_target = (noise - batch.latents).detach()
|
||||
else:
|
||||
dopsd_normal_target = noise
|
||||
elif self.sd.prediction_type == 'v_prediction':
|
||||
# v-parameterization training
|
||||
target = self.sd.noise_scheduler.get_velocity(batch.tensor, noise, timesteps)
|
||||
@@ -945,6 +958,17 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
timestep_weight = timestep_weight.view(-1, 1, 1, 1, 1).detach()
|
||||
loss = loss * timestep_weight
|
||||
|
||||
if dopsd_normal_target is not None:
|
||||
if self.train_config.loss_type == "mae":
|
||||
bleed_loss = torch.nn.functional.l1_loss(pred.float(), dopsd_normal_target.float(), reduction="none")
|
||||
else:
|
||||
bleed_loss = torch.nn.functional.mse_loss(pred.float(), dopsd_normal_target.float(), reduction="none")
|
||||
# scale normal loss to the dopsd loss magnitude, then apply bleed strength
|
||||
with torch.no_grad():
|
||||
bleed_scale = loss.detach().mean() / bleed_loss.detach().mean().clamp(min=1e-8)
|
||||
bleed_strength = float(getattr(self.sd, 'dopsd_bleed_strength', 1.0))
|
||||
loss = loss + bleed_loss * bleed_scale * bleed_strength
|
||||
|
||||
if self.train_config.do_prior_divergence and prior_pred is not None:
|
||||
loss = loss + (torch.nn.functional.mse_loss(pred.float(), prior_pred.float(), reduction="none") * -1.0)
|
||||
|
||||
|
||||
@@ -132,6 +132,8 @@ MODEL_TESTS = {
|
||||
"model": {"name_or_path": "HiDream-ai/HiDream-E1-1", "quantize": True, "quantize_te": True},
|
||||
"sample": {"width": 768, "height": 768, "num_inference_steps": 28, "guidance_scale": 5.0, "seed": 42},
|
||||
"needs_control_image": True,
|
||||
# native editing seq assert: other resolutions hard-fail
|
||||
"size_locked": True,
|
||||
},
|
||||
"nucleus_image": {
|
||||
"model": {"name_or_path": "NucleusAI/Nucleus-Image", "quantize": True, "quantize_te": True},
|
||||
@@ -219,9 +221,18 @@ def classify_error(err: BaseException) -> str:
|
||||
return "FAIL"
|
||||
|
||||
|
||||
def run_one(arch: str, device: str, allow_download: bool) -> dict:
|
||||
def run_one(
|
||||
arch: str,
|
||||
device: str,
|
||||
allow_download: bool,
|
||||
qtype_override: str = None,
|
||||
quant_only: bool = False,
|
||||
no_sample: bool = False,
|
||||
) -> dict:
|
||||
entry = MODEL_TESTS[arch]
|
||||
out_dir = os.path.join(OUTPUT_ROOT, arch.replace("/", "_").replace(":", "_"))
|
||||
if qtype_override:
|
||||
out_dir = os.path.join(out_dir, f"qtype_{qtype_override}")
|
||||
os.makedirs(out_dir, exist_ok=True)
|
||||
for old in glob.glob(os.path.join(out_dir, "*")):
|
||||
os.remove(old)
|
||||
@@ -231,6 +242,13 @@ def run_one(arch: str, device: str, allow_download: bool) -> dict:
|
||||
|
||||
# quantization always tests the convrot8 backend (per-entry override wins)
|
||||
model_kwargs = dict(entry["model"])
|
||||
if qtype_override:
|
||||
# quant smoke: force quantization of transformer + TE at this qtype
|
||||
model_kwargs["quantize"] = True
|
||||
model_kwargs["quantize_te"] = True
|
||||
model_kwargs["qtype"] = qtype_override
|
||||
model_kwargs["qtype_te"] = qtype_override
|
||||
else:
|
||||
if model_kwargs.get("quantize"):
|
||||
model_kwargs.setdefault("qtype", "convrot8")
|
||||
if model_kwargs.get("quantize_te"):
|
||||
@@ -307,6 +325,17 @@ def run_one(arch: str, device: str, allow_download: bool) -> dict:
|
||||
load_profile.dump_stats(os.path.join(out_dir, "load.prof"))
|
||||
|
||||
sample_kwargs = dict(entry["sample"])
|
||||
if quant_only:
|
||||
# quant smoke: one tiny pass just to prove the quantized forward runs.
|
||||
# Kernel-shape bugs live in layer dims (K/N), not token count, so a
|
||||
# small image exercises them; 2 steps not 1 (shift_terminal NaNs).
|
||||
sample_kwargs["num_inference_steps"] = 2
|
||||
if not entry.get("size_locked"):
|
||||
# 384 divides every bucket size in the registry (16/32/64)
|
||||
sample_kwargs["width"] = min(sample_kwargs["width"], 384)
|
||||
sample_kwargs["height"] = min(sample_kwargs["height"], 384)
|
||||
if "num_frames" in sample_kwargs:
|
||||
sample_kwargs["num_frames"] = min(sample_kwargs["num_frames"], 9)
|
||||
if entry.get("needs_control_image"):
|
||||
# edit/kontext models require a control image; a flat gray input is fine
|
||||
from PIL import Image
|
||||
@@ -329,6 +358,16 @@ def run_one(arch: str, device: str, allow_download: bool) -> dict:
|
||||
# the legacy monolith takes the sampler NAME at generate time
|
||||
gen_kwargs["sampler"] = "ddpm"
|
||||
|
||||
if quant_only and no_sample:
|
||||
# load+quantize only — no forward at all; catches quantize-time
|
||||
# errors but NOT broken quantized kernels (those need the sample)
|
||||
stages = prof.finish()
|
||||
return _finish_result(
|
||||
arch, out_dir, model_config, model_kwargs, load_seconds,
|
||||
None, None, 0, stages, load_profile, None, None,
|
||||
require_output=False,
|
||||
)
|
||||
|
||||
prof.stage("generate")
|
||||
gen_profile = cProfile.Profile()
|
||||
gen_profile.enable()
|
||||
@@ -344,6 +383,17 @@ def run_one(arch: str, device: str, allow_download: bool) -> dict:
|
||||
# for every arch ----
|
||||
import torch
|
||||
|
||||
offload_seconds = None
|
||||
attached = 0
|
||||
offload_profile = None
|
||||
if quant_only:
|
||||
stages = prof.finish()
|
||||
return _finish_result(
|
||||
arch, out_dir, model_config, model_kwargs, load_seconds,
|
||||
gen_seconds, offload_seconds, attached, stages,
|
||||
load_profile, gen_profile, offload_profile,
|
||||
)
|
||||
|
||||
from toolkit.memory_management import MemoryManager
|
||||
|
||||
prof.stage("offload 100%: attach")
|
||||
@@ -404,6 +454,21 @@ def run_one(arch: str, device: str, allow_download: bool) -> dict:
|
||||
offload_profile.dump_stats(os.path.join(out_dir, "offload.prof"))
|
||||
|
||||
stages = prof.finish()
|
||||
return _finish_result(
|
||||
arch, out_dir, model_config, model_kwargs, load_seconds,
|
||||
gen_seconds, offload_seconds, attached, stages,
|
||||
load_profile, gen_profile, offload_profile,
|
||||
)
|
||||
|
||||
|
||||
def _finish_result(
|
||||
arch, out_dir, model_config, model_kwargs, load_seconds, gen_seconds,
|
||||
offload_seconds, attached, stages, load_profile, gen_profile,
|
||||
offload_profile, require_output=True,
|
||||
):
|
||||
import torch
|
||||
|
||||
from testing.stage_profiler import profile_top
|
||||
|
||||
produced = [
|
||||
p
|
||||
@@ -412,9 +477,19 @@ def run_one(arch: str, device: str, allow_download: bool) -> dict:
|
||||
and os.path.getsize(p) > 1024
|
||||
and not p.endswith((".txt", ".prof", ".json"))
|
||||
]
|
||||
if not produced:
|
||||
if not produced and require_output:
|
||||
raise RuntimeError(f"no output file produced in {out_dir}")
|
||||
|
||||
profile = {
|
||||
"load_top": profile_top(load_profile),
|
||||
"load_prof": os.path.join(out_dir, "load.prof"),
|
||||
}
|
||||
if gen_profile is not None:
|
||||
profile["gen_top"] = profile_top(gen_profile, limit=15)
|
||||
profile["gen_prof"] = os.path.join(out_dir, "gen.prof")
|
||||
if offload_profile is not None:
|
||||
profile["offload_top"] = profile_top(offload_profile, limit=15)
|
||||
profile["offload_prof"] = os.path.join(out_dir, "offload.prof")
|
||||
result = {
|
||||
"arch": arch,
|
||||
"status": "PASS",
|
||||
@@ -422,18 +497,13 @@ def run_one(arch: str, device: str, allow_download: bool) -> dict:
|
||||
"qtype": model_config.qtype if model_kwargs.get("quantize") else None,
|
||||
"qtype_te": model_config.qtype_te if model_kwargs.get("quantize_te") else None,
|
||||
"load_seconds": round(load_seconds, 3),
|
||||
"gen_seconds": round(gen_seconds, 3),
|
||||
"offload_seconds": round(offload_seconds, 3),
|
||||
"gen_seconds": round(gen_seconds, 3) if gen_seconds is not None else None,
|
||||
"offload_seconds": (
|
||||
round(offload_seconds, 3) if offload_seconds is not None else None
|
||||
),
|
||||
"offload_modules_attached": attached,
|
||||
"stages": stages,
|
||||
"profile": {
|
||||
"load_top": profile_top(load_profile),
|
||||
"gen_top": profile_top(gen_profile, limit=15),
|
||||
"offload_top": profile_top(offload_profile, limit=15),
|
||||
"load_prof": os.path.join(out_dir, "load.prof"),
|
||||
"gen_prof": os.path.join(out_dir, "gen.prof"),
|
||||
"offload_prof": os.path.join(out_dir, "offload.prof"),
|
||||
},
|
||||
"profile": profile,
|
||||
"env": {
|
||||
"torch_num_threads": torch.get_num_threads(),
|
||||
"cpu_count": os.cpu_count(),
|
||||
@@ -574,6 +644,27 @@ def main():
|
||||
parser.add_argument("--device", type=str, default="cuda:0")
|
||||
parser.add_argument("--allow-download", action="store_true")
|
||||
parser.add_argument("--json-result", type=str, default=None)
|
||||
parser.add_argument(
|
||||
"--quant-test",
|
||||
action="store_true",
|
||||
help="quant smoke: for each arch (or --arch), force-quantize the "
|
||||
"transformer + TE at each qtype in --qtypes, run one tiny 2-step "
|
||||
"pass, no offload probe — just proves each backend loads/quantizes/"
|
||||
"runs without errors",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--qtypes",
|
||||
type=str,
|
||||
default="convrot8,qfloat8,float8",
|
||||
help="comma-separated qtypes for --quant-test",
|
||||
)
|
||||
parser.add_argument("--qtype-override", type=str, default=None)
|
||||
parser.add_argument(
|
||||
"--no-sample",
|
||||
action="store_true",
|
||||
help="with --quant-test: load+quantize only, skip the 2-step sample "
|
||||
"(faster, but does not exercise the quantized kernels)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--report-only",
|
||||
action="store_true",
|
||||
@@ -608,13 +699,73 @@ def main():
|
||||
if not args.allow_download:
|
||||
os.environ.setdefault("HF_HUB_OFFLINE", "1")
|
||||
|
||||
if args.quant_test and args.qtype_override is None:
|
||||
# quant smoke: (archs x qtypes) grid, one subprocess per cell so each
|
||||
# backend gets a clean CUDA context
|
||||
qtypes = [q.strip() for q in args.qtypes.split(",") if q.strip()]
|
||||
archs = [args.arch] if args.arch else list(MODEL_TESTS)
|
||||
results = []
|
||||
total = len(archs) * len(qtypes)
|
||||
i = 0
|
||||
t0 = time.perf_counter()
|
||||
for arch in archs:
|
||||
for qt in qtypes:
|
||||
i += 1
|
||||
print(f"\n===== {arch} [{qt}] ({i}/{total}) =====", flush=True)
|
||||
result_path = os.path.join(
|
||||
OUTPUT_ROOT, f".{arch.replace('/', '_')}.{qt}.result.json"
|
||||
)
|
||||
cmd = [
|
||||
sys.executable, os.path.abspath(__file__),
|
||||
"--arch", arch, "--device", args.device,
|
||||
"--qtype-override", qt, "--quant-test",
|
||||
"--json-result", result_path,
|
||||
]
|
||||
if args.allow_download:
|
||||
cmd.append("--allow-download")
|
||||
if args.no_sample:
|
||||
cmd.append("--no-sample")
|
||||
subprocess.run(cmd, cwd=TOOLKIT_ROOT)
|
||||
if os.path.exists(result_path):
|
||||
with open(result_path) as f:
|
||||
r = json.load(f)
|
||||
os.remove(result_path)
|
||||
else:
|
||||
r = {"arch": arch, "status": "FAIL", "error": "subprocess died"}
|
||||
r["qtype_tested"] = qt
|
||||
results.append(r)
|
||||
elapsed = time.perf_counter() - t0
|
||||
print(
|
||||
f">>> quant-test progress: {i}/{total} — elapsed "
|
||||
f"{elapsed / 60:.0f}m, eta ~{(elapsed / i) * (total - i) / 60:.0f}m",
|
||||
flush=True,
|
||||
)
|
||||
print("\n===== quant-test summary =====")
|
||||
fails = 0
|
||||
for r in results:
|
||||
line = f"[{r['status']}] {r['arch']} [{r['qtype_tested']}]"
|
||||
if r["status"] != "PASS":
|
||||
line += f" — {r.get('error', '')[:140]}"
|
||||
fails += r["status"] == "FAIL"
|
||||
print(line)
|
||||
if fails:
|
||||
sys.exit(1)
|
||||
return
|
||||
|
||||
if args.arch is not None:
|
||||
if args.arch not in MODEL_TESTS:
|
||||
raise SystemExit(
|
||||
f"arch {args.arch!r} is not registered; --list shows options"
|
||||
)
|
||||
try:
|
||||
result = run_one(args.arch, args.device, args.allow_download)
|
||||
result = run_one(
|
||||
args.arch,
|
||||
args.device,
|
||||
args.allow_download,
|
||||
qtype_override=args.qtype_override,
|
||||
quant_only=args.quant_test,
|
||||
no_sample=args.no_sample,
|
||||
)
|
||||
except BaseException as err:
|
||||
status = classify_error(err)
|
||||
result = {"arch": args.arch, "status": status, "error": f"{type(err).__name__}: {err}"}
|
||||
|
||||
@@ -146,6 +146,9 @@ class MemoryManager:
|
||||
for m in module.modules()
|
||||
if isinstance(m, torch.nn.Embedding)
|
||||
}
|
||||
# weights of embeddings that get MANAGED (cpu-resident bouncing): a
|
||||
# linear sharing one of these must be managed too, not left resident
|
||||
managed_embedding_ptrs = set()
|
||||
# attach to all modules
|
||||
for name, sub_module in module.named_modules():
|
||||
for child_name, child_module in sub_module.named_modules():
|
||||
@@ -162,6 +165,8 @@ class MemoryManager:
|
||||
not getattr(child_module, "is_ostris_quantized", False)
|
||||
and isinstance(getattr(child_module, "weight", None), torch.Tensor)
|
||||
and child_module.weight.data_ptr() in embedding_weight_ptrs
|
||||
and child_module.weight.data_ptr()
|
||||
not in managed_embedding_ptrs
|
||||
):
|
||||
skip = True
|
||||
if skip:
|
||||
@@ -226,10 +231,11 @@ class MemoryManager:
|
||||
EmbeddingLayerMemoryManager.attach(
|
||||
child_module, module._memory_manager
|
||||
)
|
||||
# the cpu move replaced weight.data, so a tied lm_head no
|
||||
# longer matches the (original-ptr) tied guard and gets
|
||||
# managed normally — correct: the shared weight now lives
|
||||
# on cpu, so the linear must bounce it
|
||||
# a tied lm_head must bounce the (now cpu-resident) shared
|
||||
# weight rather than stay resident; record both the pre-
|
||||
# and post-move ptrs (a cpu->cpu move keeps the tensor)
|
||||
managed_embedding_ptrs.add(child_module.weight.data_ptr())
|
||||
embedding_weight_ptrs.add(child_module.weight.data_ptr())
|
||||
modules_processed.append(child_module)
|
||||
elif child_module.__class__.__name__ in UNMANAGED_MODULES or any(
|
||||
inc in child_module.__class__.__name__
|
||||
@@ -249,6 +255,22 @@ class MemoryManager:
|
||||
else:
|
||||
continue
|
||||
|
||||
# everything NOT managed is the resident set and must live on the
|
||||
# compute device. A model attached while parked on cpu (the offload
|
||||
# load flow) otherwise keeps its rotary buffers / norms / conv towers
|
||||
# on cpu and the first forward explodes on a device mismatch. Managed
|
||||
# layers (pinned-cpu weights, cpu-resident bouncing embeddings) are
|
||||
# skipped via their _layer_memory_manager.
|
||||
for sub in module.modules():
|
||||
if hasattr(sub, "_layer_memory_manager"):
|
||||
continue
|
||||
for p in sub.parameters(recurse=False):
|
||||
if p is not None and p.device != device:
|
||||
p.data = p.data.to(device)
|
||||
for name, b in sub._buffers.items():
|
||||
if b is not None and b.device != device:
|
||||
sub._buffers[name] = b.to(device)
|
||||
|
||||
@classmethod
|
||||
def detach(cls, module: torch.nn.Module):
|
||||
"""
|
||||
|
||||
@@ -185,6 +185,8 @@ class BaseModel:
|
||||
self.supports_video_control_images = False
|
||||
# D-OPSD: cache per-item teacher text embeds (item's own media as reference 1)
|
||||
self.dopsd_self_ref = False
|
||||
# weight of the normal-target loss added alongside the D-OPSD teacher loss
|
||||
self.dopsd_bleed_strength = 1.0
|
||||
# forces cache_tensors_to_disk on latent-caching datasets (BaseSDTrainProcess)
|
||||
self.require_pixel_tensor_cache = False
|
||||
# control images will come in as a list for encoding some things if true
|
||||
|
||||
@@ -218,6 +218,11 @@ class OstrisModelMixin:
|
||||
config=config,
|
||||
subfolder=subfolder,
|
||||
use_comfy_weights=use_comfy_weights,
|
||||
# comfy candidate ranking prefers the file whose shipped
|
||||
# quantization matches the request (strip any ARA suffix);
|
||||
# quantization itself happens in aitk_post_load below
|
||||
qtype=(qtype or "").split("|", 1)[0] or None,
|
||||
quantize_on_load=False,
|
||||
**kwargs,
|
||||
)
|
||||
return model.aitk_post_load(
|
||||
@@ -254,6 +259,53 @@ class OstrisModelMixin:
|
||||
if qtype is not None and "|" in qtype:
|
||||
qtype, ara_path = qtype.split("|", 1)
|
||||
|
||||
# pre-quantized checkpoint: keep the shipped quantization IFF it
|
||||
# exactly matches the request. Anything else — no quantization
|
||||
# requested (full finetuning), a different backend, or an ARA —
|
||||
# restores/requantizes to what was asked for.
|
||||
if getattr(self, "aitk_is_quantized", False):
|
||||
from toolkit.util.ostris_quant import OstrisLinear
|
||||
from toolkit.util.quantize import (
|
||||
dequantize_ostris_to_linear,
|
||||
get_qtype,
|
||||
ostristype,
|
||||
)
|
||||
|
||||
shipped = sorted(
|
||||
{
|
||||
getattr(m.ostris_quantizer, "qtype", None)
|
||||
for m in self.modules()
|
||||
if isinstance(m, OstrisLinear)
|
||||
}
|
||||
- {None}
|
||||
)
|
||||
matches = (
|
||||
qtype is not None
|
||||
and ara_path is None
|
||||
and shipped == [qtype]
|
||||
)
|
||||
if matches:
|
||||
status_fn(f"Checkpoint is pre-quantized ({qtype}); keeping it")
|
||||
else:
|
||||
target_is_ostris = qtype is not None and isinstance(
|
||||
get_qtype(qtype), ostristype
|
||||
)
|
||||
if qtype is None or ara_path is not None or not target_is_ostris:
|
||||
# full precision requested, an ARA (quantizes fresh from
|
||||
# full weights), or a quanto/torchao backend that cannot
|
||||
# re-quantize an OstrisLinear: dequantize first
|
||||
status_fn(
|
||||
f"Dequantizing shipped {'/'.join(shipped)} weights "
|
||||
f"-> {qtype or 'full precision'}"
|
||||
)
|
||||
dequantize_ostris_to_linear(self)
|
||||
else:
|
||||
# ostris -> ostris: quantize_module re-quantizes each
|
||||
# layer in place (dequant -> requant, one layer at a time)
|
||||
status_fn(f"Requantizing {'/'.join(shipped)} -> {qtype}")
|
||||
self.aitk_is_quantized = False
|
||||
self.aitk_qtype = None
|
||||
|
||||
if qtype and not getattr(self, "aitk_is_quantized", False):
|
||||
from toolkit.util.quantize import (
|
||||
attach_ara_and_quantize,
|
||||
@@ -308,8 +360,6 @@ class OstrisModelMixin:
|
||||
)
|
||||
self.aitk_is_quantized = True
|
||||
self.aitk_qtype = qtype
|
||||
elif qtype and getattr(self, "aitk_is_quantized", False):
|
||||
status_fn("Checkpoint is pre-quantized; skipping quantization")
|
||||
|
||||
if offload and offload > 0:
|
||||
from toolkit.memory_management import MemoryManager
|
||||
@@ -337,6 +387,7 @@ class OstrisModelMixin:
|
||||
config=None,
|
||||
subfolder: Optional[str] = None,
|
||||
use_comfy_weights: bool = True,
|
||||
quantize_on_load: bool = True,
|
||||
**kwargs,
|
||||
):
|
||||
"""Load a model universally from a given name or path.
|
||||
@@ -370,7 +421,9 @@ class OstrisModelMixin:
|
||||
and not name_or_path.endswith(".safetensors")
|
||||
and not os.path.exists(name_or_path)
|
||||
):
|
||||
comfy_path = cls.resolve_comfy_weights(name_or_path, subfolder=subfolder)
|
||||
comfy_path = cls.resolve_comfy_weights(
|
||||
name_or_path, subfolder=subfolder, qtype=qtype
|
||||
)
|
||||
if comfy_path is not None:
|
||||
if config_path is None:
|
||||
# the standard repo supplies the config for the comfy file
|
||||
@@ -399,7 +452,9 @@ class OstrisModelMixin:
|
||||
name_or_path, subfolder=subfolder, dtype=dtype, **kwargs
|
||||
)
|
||||
|
||||
if qtype is not None:
|
||||
# quantize_on_load=False: qtype was only a comfy-candidate ranking
|
||||
# hint (the .load() path quantizes in aitk_post_load instead)
|
||||
if qtype is not None and quantize_on_load:
|
||||
model.quantize_(
|
||||
qtype, device=quantize_device, exclude=exclude_quant_modules
|
||||
)
|
||||
@@ -589,6 +644,7 @@ class OstrisModelMixin:
|
||||
local_only: bool = False,
|
||||
hf_token: Optional[str] = None,
|
||||
status_fn: Optional[callable] = None,
|
||||
qtype: Optional[str] = None,
|
||||
) -> Optional[str]:
|
||||
"""The comfy-format weight file replacing a standard ``name_or_path``,
|
||||
or None when this class has none registered for it. Best-ranked local
|
||||
@@ -618,6 +674,7 @@ class OstrisModelMixin:
|
||||
hf_token=hf_token,
|
||||
status_fn=status_fn,
|
||||
local_only=local_only,
|
||||
qtype=qtype,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@@ -16,22 +16,41 @@ from typing import Callable, Iterable, Optional
|
||||
from toolkit.paths import MODELS_PATH
|
||||
|
||||
|
||||
def comfy_precision_rank(filename: str) -> int:
|
||||
"""Load-preference rank for a comfy weight filename:
|
||||
convrot8 (0) > float8 mixed (1) > float8 (2) > bf16 (3) > fp16 (4) >
|
||||
anything else, e.g. nvfp4 or unmarked (5)."""
|
||||
def comfy_precision_rank(filename: str, qtype: Optional[str] = None) -> int:
|
||||
"""Load-preference rank for a comfy weight filename, given the REQUESTED
|
||||
quantization. A file whose shipped quantization matches the request loads
|
||||
with no work; anything else costs a dequantize/requantize pass, so:
|
||||
|
||||
- convrot* requested: convrot (0) > fp8 mixed (1) > fp8 (2) > bf16 (3) >
|
||||
fp16 (4) > other (5)
|
||||
- float8/qfloat8 requested: fp8 mixed (0) > fp8 (1) > bf16 (2) > fp16 (3)
|
||||
> convrot (4) > other (5)
|
||||
- nvfp4 requested: nvfp4 (0) > bf16 (1) > fp16 (2) > convrot (3) > fp8
|
||||
(4) > other (5)
|
||||
- no quantization requested (full precision) or any other fresh-quant
|
||||
backend: bf16 (0) > fp16 (1) > convrot (2) > fp8 mixed (3) > fp8 (4) >
|
||||
other (5) — clean weights beat paying a dequantize
|
||||
"""
|
||||
name = os.path.basename(filename).lower()
|
||||
if "convrot" in name:
|
||||
return 0
|
||||
is_fp8 = "fp8" in name or "float8" in name or "e4m3" in name
|
||||
if is_fp8 and "mixed" in name:
|
||||
return 1
|
||||
if is_fp8:
|
||||
return 2
|
||||
if "bf16" in name:
|
||||
return 3
|
||||
if "fp16" in name:
|
||||
return 4
|
||||
is_convrot = "convrot" in name
|
||||
is_nvfp4 = "nvfp4" in name
|
||||
is_fp8 = ("fp8" in name or "float8" in name or "e4m3" in name) and not is_nvfp4
|
||||
is_fp8_mixed = is_fp8 and "mixed" in name
|
||||
is_bf16 = "bf16" in name
|
||||
is_fp16 = "fp16" in name and not is_fp8
|
||||
|
||||
qt = (qtype or "").lower()
|
||||
if qt.startswith("convrot"):
|
||||
order = [is_convrot, is_fp8_mixed, is_fp8, is_bf16, is_fp16]
|
||||
elif "float8" in qt or qt == "qfloat8":
|
||||
order = [is_fp8_mixed, is_fp8, is_bf16, is_fp16, is_convrot]
|
||||
elif "nvfp4" in qt:
|
||||
order = [is_nvfp4, is_bf16, is_fp16, is_convrot, is_fp8]
|
||||
else:
|
||||
order = [is_bf16, is_fp16, is_convrot, is_fp8_mixed, is_fp8]
|
||||
for rank, flag in enumerate(order):
|
||||
if flag:
|
||||
return rank
|
||||
return 5
|
||||
|
||||
|
||||
@@ -51,15 +70,17 @@ def resolve_comfy_candidates(
|
||||
hf_token: Optional[str] = None,
|
||||
status_fn: Optional[Callable[[str], None]] = None,
|
||||
local_only: bool = False,
|
||||
qtype: Optional[str] = None,
|
||||
) -> Optional[str]:
|
||||
"""Pick the best comfy weight file among precision variants of one
|
||||
component (repo-relative paths, ranked by comfy_precision_rank then list
|
||||
order). The best-ranked LOCAL candidate wins; only when no candidate is
|
||||
local is the best-ranked one downloaded to its comfy-layout location
|
||||
under MODELS_PATH."""
|
||||
component (repo-relative paths, ranked by comfy_precision_rank for the
|
||||
requested qtype, then list order). The best-ranked LOCAL candidate wins;
|
||||
only when no candidate is local is the best-ranked one downloaded to its
|
||||
comfy-layout location under MODELS_PATH."""
|
||||
candidates = list(candidates)
|
||||
ordered = sorted(
|
||||
candidates, key=lambda c: (comfy_precision_rank(c), candidates.index(c))
|
||||
candidates,
|
||||
key=lambda c: (comfy_precision_rank(c, qtype=qtype), candidates.index(c)),
|
||||
)
|
||||
for repo_rel in ordered:
|
||||
found = resolve_comfy_file(
|
||||
|
||||
@@ -260,6 +260,8 @@ class StableDiffusion:
|
||||
self.supports_video_control_images = False
|
||||
# D-OPSD: cache per-item teacher text embeds (item's own media as reference 1)
|
||||
self.dopsd_self_ref = False
|
||||
# weight of the normal-target loss added alongside the D-OPSD teacher loss
|
||||
self.dopsd_bleed_strength = 1.0
|
||||
# forces cache_tensors_to_disk on latent-caching datasets (BaseSDTrainProcess)
|
||||
self.require_pixel_tensor_cache = False
|
||||
# control images will come in as a list for encoding some things if true
|
||||
|
||||
@@ -125,6 +125,22 @@ def requantize_module_weight(module, fp_weight, orig_dtype, config) -> None:
|
||||
torchao_quantize_(module, config)
|
||||
|
||||
|
||||
def _wrap_qlinear_ndim(qlinear: torch.nn.Module) -> None:
|
||||
"""quanto QLinear forward that tolerates >3D activations by flattening
|
||||
the leading dims for the mm and restoring them after (no-op otherwise)."""
|
||||
orig_forward = qlinear.forward
|
||||
|
||||
def forward(x):
|
||||
if x.ndim > 3:
|
||||
lead = x.shape[:-1]
|
||||
out = orig_forward(x.reshape(-1, x.shape[-1]))
|
||||
return out.reshape(*lead, out.shape[-1])
|
||||
return orig_forward(x)
|
||||
|
||||
qlinear.forward = forward
|
||||
qlinear._aitk_ndim_wrapped = True
|
||||
|
||||
|
||||
def quantize(
|
||||
model: torch.nn.Module,
|
||||
weights: Optional[Union[str, qtype, aotype]] = None,
|
||||
@@ -242,6 +258,14 @@ def quantize(
|
||||
activations=activations,
|
||||
optimizer=optimizer,
|
||||
)
|
||||
# quanto's qbytes_mm only takes 2D/3D activations; video
|
||||
# patch embeds feed their linears >3D tensors (convrot/
|
||||
# torchao reshape internally). Flatten around the QLinear.
|
||||
replaced = model.get_submodule(name)
|
||||
if replaced.__class__.__name__ == "QLinear" and not getattr(
|
||||
replaced, "_aitk_ndim_wrapped", False
|
||||
):
|
||||
_wrap_qlinear_ndim(replaced)
|
||||
finally:
|
||||
if orig_device is not None and not keep_on_quantize_device:
|
||||
# quanto replaces the module in its parent, so re-fetch by name
|
||||
@@ -279,6 +303,39 @@ def _has_quantizable_linear(module: torch.nn.Module, weights, exclude=None) -> b
|
||||
return False
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def dequantize_ostris_to_linear(module: torch.nn.Module) -> int:
|
||||
"""Replace every OstrisLinear with a plain nn.Linear holding the full-
|
||||
precision weight (activation-side transforms folded), in place, layer by
|
||||
layer — the full-precision transient never exceeds one layer. Used when a
|
||||
pre-quantized checkpoint is loaded but a DIFFERENT quantization (or none,
|
||||
e.g. full finetuning) was requested. Returns the number of layers
|
||||
restored."""
|
||||
replaced = 0
|
||||
for parent in module.modules():
|
||||
for child_name, child in list(parent.named_children()):
|
||||
if not isinstance(child, OstrisLinear):
|
||||
continue
|
||||
weight = child.ostris_quantizer.dequantize_folded(child).to(
|
||||
child.ostris_orig_dtype
|
||||
)
|
||||
new = torch.nn.Linear(
|
||||
child.in_features,
|
||||
child.out_features,
|
||||
bias=child.bias is not None,
|
||||
device="meta",
|
||||
dtype=weight.dtype,
|
||||
)
|
||||
new.weight = torch.nn.Parameter(weight)
|
||||
if child.bias is not None:
|
||||
new.bias = torch.nn.Parameter(
|
||||
child.bias.data.to(weight.device, weight.dtype)
|
||||
)
|
||||
setattr(parent, child_name, new)
|
||||
replaced += 1
|
||||
return replaced
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def quantize_module(
|
||||
module: torch.nn.Module,
|
||||
|
||||
@@ -967,8 +967,10 @@ export const modelArchs: ModelArch[] = [
|
||||
const kwargs = { ...(config?.config?.process?.[0]?.model?.model_kwargs ?? {}) };
|
||||
if (value === 'dopsd') {
|
||||
kwargs.dopsd = true;
|
||||
kwargs.dopsd_bleed_strength = 1.0;
|
||||
} else {
|
||||
delete kwargs.dopsd;
|
||||
delete kwargs.dopsd_bleed_strength;
|
||||
}
|
||||
setJobConfig(kwargs, 'config.process[0].model.model_kwargs');
|
||||
if (value === 'cg') {
|
||||
|
||||
Reference in New Issue
Block a user