Stability improvements to model offloading. Added D-OPSD bleed loss as well.

This commit is contained in:
Jaret Burkett
2026-08-29 11:23:47 -06:00
parent 64a20f51a6
commit 683fe8afc0
11 changed files with 390 additions and 45 deletions

View File

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

View File

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

View File

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

View File

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

View 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

View File

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