Stability improvements to model offloading. Added D-OPSD bleed loss as well.
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user