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

@@ -1153,6 +1153,9 @@ class MinimaxH3Ref2VAModel(MinimaxH3Model):
if self.dopsd: if self.dopsd:
self.dopsd_self_ref = True self.dopsd_self_ref = True
self.require_pixel_tensor_cache = 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: def _dit_component(self) -> str:
partition = str( partition = str(

View File

@@ -188,6 +188,10 @@ class ZImageModel(BaseModel):
use_comfy_weights=self.model_config.model_kwargs.get( use_comfy_weights=self.model_config.model_kwargs.get(
"use_comfy_weights", True "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() flush()

View File

@@ -556,6 +556,7 @@ class SDTrainer(BaseSDTrainProcess):
noise_pred = noise_pred * self.train_config.pred_scaler noise_pred = noise_pred * self.train_config.pred_scaler
target = None target = None
dopsd_normal_target = None
if self.train_config.target_noise_multiplier != 1.0: if self.train_config.target_noise_multiplier != 1.0:
noise = noise * self.train_config.target_noise_multiplier noise = noise * self.train_config.target_noise_multiplier
@@ -623,6 +624,18 @@ class SDTrainer(BaseSDTrainProcess):
assert not self.train_config.train_turbo assert not self.train_config.train_turbo
# matching adapter prediction # matching adapter prediction
target = prior_pred 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': elif self.sd.prediction_type == 'v_prediction':
# v-parameterization training # v-parameterization training
target = self.sd.noise_scheduler.get_velocity(batch.tensor, noise, timesteps) 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() timestep_weight = timestep_weight.view(-1, 1, 1, 1, 1).detach()
loss = loss * timestep_weight 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: 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) loss = loss + (torch.nn.functional.mse_loss(pred.float(), prior_pred.float(), reduction="none") * -1.0)

View File

@@ -132,6 +132,8 @@ MODEL_TESTS = {
"model": {"name_or_path": "HiDream-ai/HiDream-E1-1", "quantize": True, "quantize_te": True}, "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}, "sample": {"width": 768, "height": 768, "num_inference_steps": 28, "guidance_scale": 5.0, "seed": 42},
"needs_control_image": True, "needs_control_image": True,
# native editing seq assert: other resolutions hard-fail
"size_locked": True,
}, },
"nucleus_image": { "nucleus_image": {
"model": {"name_or_path": "NucleusAI/Nucleus-Image", "quantize": True, "quantize_te": True}, "model": {"name_or_path": "NucleusAI/Nucleus-Image", "quantize": True, "quantize_te": True},
@@ -219,9 +221,18 @@ def classify_error(err: BaseException) -> str:
return "FAIL" 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] entry = MODEL_TESTS[arch]
out_dir = os.path.join(OUTPUT_ROOT, arch.replace("/", "_").replace(":", "_")) 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) os.makedirs(out_dir, exist_ok=True)
for old in glob.glob(os.path.join(out_dir, "*")): for old in glob.glob(os.path.join(out_dir, "*")):
os.remove(old) os.remove(old)
@@ -231,10 +242,17 @@ def run_one(arch: str, device: str, allow_download: bool) -> dict:
# quantization always tests the convrot8 backend (per-entry override wins) # quantization always tests the convrot8 backend (per-entry override wins)
model_kwargs = dict(entry["model"]) model_kwargs = dict(entry["model"])
if model_kwargs.get("quantize"): if qtype_override:
model_kwargs.setdefault("qtype", "convrot8") # quant smoke: force quantization of transformer + TE at this qtype
if model_kwargs.get("quantize_te"): model_kwargs["quantize"] = True
model_kwargs.setdefault("qtype_te", "convrot8") 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"):
model_kwargs.setdefault("qtype_te", "convrot8")
model_config = ModelConfig(arch=arch, dtype="bf16", **model_kwargs) model_config = ModelConfig(arch=arch, dtype="bf16", **model_kwargs)
ModelClass = get_model_class(model_config) ModelClass = get_model_class(model_config)
@@ -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")) load_profile.dump_stats(os.path.join(out_dir, "load.prof"))
sample_kwargs = dict(entry["sample"]) 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"): if entry.get("needs_control_image"):
# edit/kontext models require a control image; a flat gray input is fine # edit/kontext models require a control image; a flat gray input is fine
from PIL import Image 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 # the legacy monolith takes the sampler NAME at generate time
gen_kwargs["sampler"] = "ddpm" 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") prof.stage("generate")
gen_profile = cProfile.Profile() gen_profile = cProfile.Profile()
gen_profile.enable() gen_profile.enable()
@@ -344,6 +383,17 @@ def run_one(arch: str, device: str, allow_download: bool) -> dict:
# for every arch ---- # for every arch ----
import torch 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 from toolkit.memory_management import MemoryManager
prof.stage("offload 100%: attach") 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")) offload_profile.dump_stats(os.path.join(out_dir, "offload.prof"))
stages = prof.finish() 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 = [ produced = [
p p
@@ -412,9 +477,19 @@ def run_one(arch: str, device: str, allow_download: bool) -> dict:
and os.path.getsize(p) > 1024 and os.path.getsize(p) > 1024
and not p.endswith((".txt", ".prof", ".json")) 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}") 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 = { result = {
"arch": arch, "arch": arch,
"status": "PASS", "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": model_config.qtype if model_kwargs.get("quantize") else None,
"qtype_te": model_config.qtype_te if model_kwargs.get("quantize_te") else None, "qtype_te": model_config.qtype_te if model_kwargs.get("quantize_te") else None,
"load_seconds": round(load_seconds, 3), "load_seconds": round(load_seconds, 3),
"gen_seconds": round(gen_seconds, 3), "gen_seconds": round(gen_seconds, 3) if gen_seconds is not None else None,
"offload_seconds": round(offload_seconds, 3), "offload_seconds": (
round(offload_seconds, 3) if offload_seconds is not None else None
),
"offload_modules_attached": attached, "offload_modules_attached": attached,
"stages": stages, "stages": stages,
"profile": { "profile": 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"),
},
"env": { "env": {
"torch_num_threads": torch.get_num_threads(), "torch_num_threads": torch.get_num_threads(),
"cpu_count": os.cpu_count(), "cpu_count": os.cpu_count(),
@@ -574,6 +644,27 @@ def main():
parser.add_argument("--device", type=str, default="cuda:0") parser.add_argument("--device", type=str, default="cuda:0")
parser.add_argument("--allow-download", action="store_true") parser.add_argument("--allow-download", action="store_true")
parser.add_argument("--json-result", type=str, default=None) 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( parser.add_argument(
"--report-only", "--report-only",
action="store_true", action="store_true",
@@ -608,13 +699,73 @@ def main():
if not args.allow_download: if not args.allow_download:
os.environ.setdefault("HF_HUB_OFFLINE", "1") 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 is not None:
if args.arch not in MODEL_TESTS: if args.arch not in MODEL_TESTS:
raise SystemExit( raise SystemExit(
f"arch {args.arch!r} is not registered; --list shows options" f"arch {args.arch!r} is not registered; --list shows options"
) )
try: 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: except BaseException as err:
status = classify_error(err) status = classify_error(err)
result = {"arch": args.arch, "status": status, "error": f"{type(err).__name__}: {err}"} result = {"arch": args.arch, "status": status, "error": f"{type(err).__name__}: {err}"}

View File

@@ -146,6 +146,9 @@ class MemoryManager:
for m in module.modules() for m in module.modules()
if isinstance(m, torch.nn.Embedding) 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 # attach to all modules
for name, sub_module in module.named_modules(): for name, sub_module in module.named_modules():
for child_name, child_module in sub_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) not getattr(child_module, "is_ostris_quantized", False)
and isinstance(getattr(child_module, "weight", None), torch.Tensor) 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() in embedding_weight_ptrs
and child_module.weight.data_ptr()
not in managed_embedding_ptrs
): ):
skip = True skip = True
if skip: if skip:
@@ -226,10 +231,11 @@ class MemoryManager:
EmbeddingLayerMemoryManager.attach( EmbeddingLayerMemoryManager.attach(
child_module, module._memory_manager child_module, module._memory_manager
) )
# the cpu move replaced weight.data, so a tied lm_head no # a tied lm_head must bounce the (now cpu-resident) shared
# longer matches the (original-ptr) tied guard and gets # weight rather than stay resident; record both the pre-
# managed normally — correct: the shared weight now lives # and post-move ptrs (a cpu->cpu move keeps the tensor)
# on cpu, so the linear must bounce it managed_embedding_ptrs.add(child_module.weight.data_ptr())
embedding_weight_ptrs.add(child_module.weight.data_ptr())
modules_processed.append(child_module) modules_processed.append(child_module)
elif child_module.__class__.__name__ in UNMANAGED_MODULES or any( elif child_module.__class__.__name__ in UNMANAGED_MODULES or any(
inc in child_module.__class__.__name__ inc in child_module.__class__.__name__
@@ -249,6 +255,22 @@ class MemoryManager:
else: else:
continue 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 @classmethod
def detach(cls, module: torch.nn.Module): def detach(cls, module: torch.nn.Module):
""" """

View File

@@ -185,6 +185,8 @@ class BaseModel:
self.supports_video_control_images = False self.supports_video_control_images = False
# D-OPSD: cache per-item teacher text embeds (item's own media as reference 1) # D-OPSD: cache per-item teacher text embeds (item's own media as reference 1)
self.dopsd_self_ref = False 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) # forces cache_tensors_to_disk on latent-caching datasets (BaseSDTrainProcess)
self.require_pixel_tensor_cache = False self.require_pixel_tensor_cache = False
# control images will come in as a list for encoding some things if true # control images will come in as a list for encoding some things if true

View File

@@ -218,6 +218,11 @@ class OstrisModelMixin:
config=config, config=config,
subfolder=subfolder, subfolder=subfolder,
use_comfy_weights=use_comfy_weights, 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, **kwargs,
) )
return model.aitk_post_load( return model.aitk_post_load(
@@ -254,6 +259,53 @@ class OstrisModelMixin:
if qtype is not None and "|" in qtype: if qtype is not None and "|" in qtype:
qtype, ara_path = qtype.split("|", 1) 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): if qtype and not getattr(self, "aitk_is_quantized", False):
from toolkit.util.quantize import ( from toolkit.util.quantize import (
attach_ara_and_quantize, attach_ara_and_quantize,
@@ -308,8 +360,6 @@ class OstrisModelMixin:
) )
self.aitk_is_quantized = True self.aitk_is_quantized = True
self.aitk_qtype = qtype 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: if offload and offload > 0:
from toolkit.memory_management import MemoryManager from toolkit.memory_management import MemoryManager
@@ -337,6 +387,7 @@ class OstrisModelMixin:
config=None, config=None,
subfolder: Optional[str] = None, subfolder: Optional[str] = None,
use_comfy_weights: bool = True, use_comfy_weights: bool = True,
quantize_on_load: bool = True,
**kwargs, **kwargs,
): ):
"""Load a model universally from a given name or path. """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 name_or_path.endswith(".safetensors")
and not os.path.exists(name_or_path) 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 comfy_path is not None:
if config_path is None: if config_path is None:
# the standard repo supplies the config for the comfy file # the standard repo supplies the config for the comfy file
@@ -399,7 +452,9 @@ class OstrisModelMixin:
name_or_path, subfolder=subfolder, dtype=dtype, **kwargs 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_( model.quantize_(
qtype, device=quantize_device, exclude=exclude_quant_modules qtype, device=quantize_device, exclude=exclude_quant_modules
) )
@@ -589,6 +644,7 @@ class OstrisModelMixin:
local_only: bool = False, local_only: bool = False,
hf_token: Optional[str] = None, hf_token: Optional[str] = None,
status_fn: Optional[callable] = None, status_fn: Optional[callable] = None,
qtype: Optional[str] = None,
) -> Optional[str]: ) -> Optional[str]:
"""The comfy-format weight file replacing a standard ``name_or_path``, """The comfy-format weight file replacing a standard ``name_or_path``,
or None when this class has none registered for it. Best-ranked local or None when this class has none registered for it. Best-ranked local
@@ -618,6 +674,7 @@ class OstrisModelMixin:
hf_token=hf_token, hf_token=hf_token,
status_fn=status_fn, status_fn=status_fn,
local_only=local_only, local_only=local_only,
qtype=qtype,
) )
# ------------------------------------------------------------------ # ------------------------------------------------------------------

View File

@@ -16,22 +16,41 @@ from typing import Callable, Iterable, Optional
from toolkit.paths import MODELS_PATH from toolkit.paths import MODELS_PATH
def comfy_precision_rank(filename: str) -> int: def comfy_precision_rank(filename: str, qtype: Optional[str] = None) -> int:
"""Load-preference rank for a comfy weight filename: """Load-preference rank for a comfy weight filename, given the REQUESTED
convrot8 (0) > float8 mixed (1) > float8 (2) > bf16 (3) > fp16 (4) > quantization. A file whose shipped quantization matches the request loads
anything else, e.g. nvfp4 or unmarked (5).""" 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() name = os.path.basename(filename).lower()
if "convrot" in name: is_convrot = "convrot" in name
return 0 is_nvfp4 = "nvfp4" in name
is_fp8 = "fp8" in name or "float8" in name or "e4m3" in name is_fp8 = ("fp8" in name or "float8" in name or "e4m3" in name) and not is_nvfp4
if is_fp8 and "mixed" in name: is_fp8_mixed = is_fp8 and "mixed" in name
return 1 is_bf16 = "bf16" in name
if is_fp8: is_fp16 = "fp16" in name and not is_fp8
return 2
if "bf16" in name: qt = (qtype or "").lower()
return 3 if qt.startswith("convrot"):
if "fp16" in name: order = [is_convrot, is_fp8_mixed, is_fp8, is_bf16, is_fp16]
return 4 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 return 5
@@ -51,15 +70,17 @@ def resolve_comfy_candidates(
hf_token: Optional[str] = None, hf_token: Optional[str] = None,
status_fn: Optional[Callable[[str], None]] = None, status_fn: Optional[Callable[[str], None]] = None,
local_only: bool = False, local_only: bool = False,
qtype: Optional[str] = None,
) -> Optional[str]: ) -> Optional[str]:
"""Pick the best comfy weight file among precision variants of one """Pick the best comfy weight file among precision variants of one
component (repo-relative paths, ranked by comfy_precision_rank then list component (repo-relative paths, ranked by comfy_precision_rank for the
order). The best-ranked LOCAL candidate wins; only when no candidate is requested qtype, then list order). The best-ranked LOCAL candidate wins;
local is the best-ranked one downloaded to its comfy-layout location only when no candidate is local is the best-ranked one downloaded to its
under MODELS_PATH.""" comfy-layout location under MODELS_PATH."""
candidates = list(candidates) candidates = list(candidates)
ordered = sorted( 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: for repo_rel in ordered:
found = resolve_comfy_file( found = resolve_comfy_file(

View File

@@ -260,6 +260,8 @@ class StableDiffusion:
self.supports_video_control_images = False self.supports_video_control_images = False
# D-OPSD: cache per-item teacher text embeds (item's own media as reference 1) # D-OPSD: cache per-item teacher text embeds (item's own media as reference 1)
self.dopsd_self_ref = False 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) # forces cache_tensors_to_disk on latent-caching datasets (BaseSDTrainProcess)
self.require_pixel_tensor_cache = False self.require_pixel_tensor_cache = False
# control images will come in as a list for encoding some things if true # 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) 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( def quantize(
model: torch.nn.Module, model: torch.nn.Module,
weights: Optional[Union[str, qtype, aotype]] = None, weights: Optional[Union[str, qtype, aotype]] = None,
@@ -242,6 +258,14 @@ def quantize(
activations=activations, activations=activations,
optimizer=optimizer, 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: finally:
if orig_device is not None and not keep_on_quantize_device: if orig_device is not None and not keep_on_quantize_device:
# quanto replaces the module in its parent, so re-fetch by name # 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 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() @torch.no_grad()
def quantize_module( def quantize_module(
module: torch.nn.Module, module: torch.nn.Module,

View File

@@ -967,8 +967,10 @@ export const modelArchs: ModelArch[] = [
const kwargs = { ...(config?.config?.process?.[0]?.model?.model_kwargs ?? {}) }; const kwargs = { ...(config?.config?.process?.[0]?.model?.model_kwargs ?? {}) };
if (value === 'dopsd') { if (value === 'dopsd') {
kwargs.dopsd = true; kwargs.dopsd = true;
kwargs.dopsd_bleed_strength = 1.0;
} else { } else {
delete kwargs.dopsd; delete kwargs.dopsd;
delete kwargs.dopsd_bleed_strength;
} }
setJobConfig(kwargs, 'config.process[0].model.model_kwargs'); setJobConfig(kwargs, 'config.process[0].model.model_kwargs');
if (value === 'cg') { if (value === 'cg') {