Fix quantization issues

This commit is contained in:
Jaret Burkett
2026-08-28 10:48:34 -06:00
parent 85a6880643
commit 92df289931
7 changed files with 470 additions and 19 deletions

View File

@@ -1,5 +1,5 @@
import os
from toolkit.models.v2.text_encoders.qwen3_vl import Qwen3VLTextEncoder
from toolkit.models.v2._mixin import OstrisTransformersMixin
from typing import List, Optional
import torch
@@ -28,6 +28,17 @@ from .src.hidream_o1.model_config import model_config
if TYPE_CHECKING:
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
class HidreamO1Transformer(Qwen3VLForConditionalGeneration, OstrisTransformersMixin):
"""The o1 DiT-in-LLM: the vendored Qwen3VL with the image-diffusion heads
(x_embedder / t_embedder1 / final_layer2). The generic Qwen3VLTextEncoder
must NOT be used here — it drops those keys as unexpected."""
@classmethod
def get_transformer_block_names(cls):
return ["model.language_model.layers"]
scheduler_config = {
"num_train_timesteps": 1000,
"shift": 3.0,
@@ -129,7 +140,10 @@ class HidreamO1Model(BaseModel):
self.use_old_lokr_format = False
self.is_flow_matching = True
self.is_transformer = True
self.target_lora_modules = ["Qwen3VLForConditionalGeneration"]
self.target_lora_modules = [
"Qwen3VLForConditionalGeneration",
"HidreamO1Transformer",
]
self.noise_scale = self.model_config.model_kwargs.get(
"noise_scale", DEFAULT_NOISE_SCALE
)
@@ -188,7 +202,7 @@ class HidreamO1Model(BaseModel):
)
# transformer.load_state_dict(state_dict, assign=True)
transformer = Qwen3VLTextEncoder.from_pretrained(
transformer = HidreamO1Transformer.from_pretrained(
None,
config=Qwen3VLConfig(**model_config),
state_dict=state_dict,
@@ -196,7 +210,7 @@ class HidreamO1Model(BaseModel):
)
del state_dict # free memory
else:
transformer = Qwen3VLTextEncoder.from_pretrained(
transformer = HidreamO1Transformer.from_pretrained(
model_path,
torch_dtype=self.torch_dtype,
)

157
testing/stage_profiler.py Normal file
View File

@@ -0,0 +1,157 @@
"""Per-stage speed / memory / CPU profiler for the model loading test harness.
A StageProfiler splits a run into named stages (the harness hooks
BaseModel.print_and_status_update so every holder status line — "Loading
transformer", "Quantizing (convrot8)", "Loading VAE", ... — opens a new
stage). For every stage it records:
- seconds wall time
- vram_peak_gb torch.cuda.max_memory_allocated during the stage
- vram_reserved_gb torch.cuda.max_memory_reserved during the stage
- vram_end_gb memory_allocated at the stage boundary (resident model)
- rss_peak_gb peak process RSS (sampled at 50ms)
- cpu_avg_pct / process CPU utilization (100 = one full core), sampled
cpu_max_pct at 50ms — this is what shows where quantization burns CPU
- threads_max peak thread count
A background sampler thread does the RSS/CPU sampling so short spikes inside
a stage are caught, not just the boundary values.
"""
import threading
import time
import psutil
import torch
_GB = 1024**3
class _Sampler:
"""Background thread sampling RSS / CPU% / thread count at 50ms."""
def __init__(self):
self.proc = psutil.Process()
self._lock = threading.Lock()
self._stop = False
self.reset()
self.proc.cpu_percent(None) # prime the cpu_percent window
self._thread = threading.Thread(target=self._run, daemon=True)
self._thread.start()
def reset(self):
with self._lock:
self.rss_peak = self.proc.memory_info().rss
self.cpu_samples = []
self.threads_max = self.proc.num_threads()
def _run(self):
while not self._stop:
try:
rss = self.proc.memory_info().rss
cpu = self.proc.cpu_percent(None)
threads = self.proc.num_threads()
with self._lock:
self.rss_peak = max(self.rss_peak, rss)
self.cpu_samples.append(cpu)
self.threads_max = max(self.threads_max, threads)
except Exception:
pass
time.sleep(0.05)
def stats(self):
with self._lock:
samples = self.cpu_samples or [0.0]
return {
"rss_peak_gb": round(self.rss_peak / _GB, 3),
"cpu_avg_pct": round(sum(samples) / len(samples), 1),
"cpu_max_pct": round(max(samples), 1),
"threads_max": self.threads_max,
}
def stop(self):
self._stop = True
class StageProfiler:
def __init__(self, device):
self.device = torch.device(device) if torch.cuda.is_available() else None
self.stages = []
self._current = None
self._sampler = _Sampler()
def _cuda_sync_reset(self):
if self.device is not None:
torch.cuda.synchronize(self.device)
torch.cuda.reset_peak_memory_stats(self.device)
def stage(self, name: str):
"""Close the current stage (if any) and open a new one."""
self._close()
self._cuda_sync_reset()
self._sampler.reset()
self._current = {"name": str(name), "t0": time.perf_counter()}
def _close(self):
if self._current is None:
return
cur = self._current
self._current = None
seconds = time.perf_counter() - cur["t0"]
entry = {"name": cur["name"], "seconds": round(seconds, 3)}
if self.device is not None:
torch.cuda.synchronize(self.device)
entry["vram_peak_gb"] = round(
torch.cuda.max_memory_allocated(self.device) / _GB, 3
)
entry["vram_reserved_gb"] = round(
torch.cuda.max_memory_reserved(self.device) / _GB, 3
)
entry["vram_end_gb"] = round(
torch.cuda.memory_allocated(self.device) / _GB, 3
)
entry.update(self._sampler.stats())
self.stages.append(entry)
def finish(self):
self._close()
self._sampler.stop()
return self.stages
def profile_top(profile, limit=40):
"""Top functions from a cProfile.Profile by cumulative time, as dicts."""
import pstats
stats = pstats.Stats(profile)
stats.sort_stats("cumulative")
rows = []
for func in stats.fcn_list[: limit * 3]:
cc, nc, tt, ct, _ = stats.stats[func]
filename, line, name = func
# drop the profiler/exec wrappers and trivial rows
if name in ("<module>",) or ct < 0.05:
continue
# cProfile catches this module's own 50ms sampling thread; its idle
# sleep loop and psutil polling otherwise show up as a giant fake
# "time.sleep" hotspot spanning the whole stage
if "stage_profiler" in filename or "psutil" in filename or "_pslinux" in filename:
continue
if name == "<built-in method time.sleep>":
continue
# shorten site-packages / repo paths to keep the report readable
for marker in ("site-packages/", "ai-toolkit/"):
if marker in filename:
filename = filename.split(marker, 1)[1]
break
rows.append(
{
"func": f"{filename}:{line}({name})",
"ncalls": nc,
"tottime": round(tt, 3),
"cumtime": round(ct, 3),
}
)
if len(rows) >= limit:
break
return rows

View File

@@ -23,6 +23,7 @@ import json
import os
import subprocess
import sys
import time
TOOLKIT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, TOOLKIT_ROOT)
@@ -57,11 +58,11 @@ MODEL_TESTS = {
"boogu_image": {
# native ~1024; 512/low-step/high-CFG degenerates to a black frame
"model": {"name_or_path": "Boogu/Boogu-Image-0.1-Base", "quantize": True, "quantize_te": True},
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
"sample": {"width": 512, "height": 512, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
},
"ernie_image": {
"model": {"name_or_path": "baidu/ERNIE-Image", "quantize": True, "quantize_te": True},
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
"sample": {"width": 512, "height": 512, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
},
"mageflow": {
"model": {"name_or_path": "microsoft/Mage-Flow-Base", "quantize": True, "quantize_te": True},
@@ -69,15 +70,15 @@ MODEL_TESTS = {
},
"ideogram4": {
"model": {"name_or_path": "ideogram-ai/ideogram-4-fp8", "quantize": True, "quantize_te": True},
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
"sample": {"width": 512, "height": 512, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
},
"hidream_o1": {
"model": {"name_or_path": "HiDream-ai/HiDream-O1-Image", "quantize": True, "quantize_te": True},
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 28, "guidance_scale": 5.0, "seed": 42},
"sample": {"width": 512, "height": 512, "num_inference_steps": 28, "guidance_scale": 5.0, "seed": 42},
},
"anima": {
"model": {"name_or_path": "circlestone-labs/Anima-Base-v1.0-Diffusers"},
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 25, "guidance_scale": 4.5, "seed": 42},
"sample": {"width": 512, "height": 512, "num_inference_steps": 25, "guidance_scale": 4.5, "seed": 42},
},
"wan21": {
"model": {"name_or_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"},
@@ -124,7 +125,7 @@ MODEL_TESTS = {
},
"hidream": {
"model": {"name_or_path": "HiDream-ai/HiDream-I1-Full", "quantize": True, "quantize_te": True},
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 28, "guidance_scale": 5.0, "seed": 42},
"sample": {"width": 512, "height": 512, "num_inference_steps": 28, "guidance_scale": 5.0, "seed": 42},
},
"hidream_e1": {
# editing seq budget fits 768x768 (source+target concat)
@@ -134,11 +135,11 @@ MODEL_TESTS = {
},
"nucleus_image": {
"model": {"name_or_path": "NucleusAI/Nucleus-Image", "quantize": True, "quantize_te": True},
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
"sample": {"width": 512, "height": 512, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
},
"omnigen2": {
"model": {"name_or_path": "OmniGen2/OmniGen2", "quantize": True, "quantize_te": True},
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
"sample": {"width": 512, "height": 512, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
},
"ltx2.5": {
"model": {"name_or_path": "Lightricks/LTX-2.5", "quantize": True, "quantize_te": True},
@@ -181,7 +182,7 @@ MODEL_TESTS = {
},
"sdxl": {
"model": {"name_or_path": "stabilityai/stable-diffusion-xl-base-1.0"},
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 25, "guidance_scale": 6.0, "seed": 42},
"sample": {"width": 512, "height": 512, "num_inference_steps": 25, "guidance_scale": 6.0, "seed": 42},
},
"ace_step_15": {
"model": {"name_or_path": "ostris/ace_step_1.5_ComfyUI_files/ace_step_1.5_base_aio.safetensors", "quantize": True, "quantize_te": True},
@@ -189,7 +190,7 @@ MODEL_TESTS = {
},
"f-lite": {
"model": {"name_or_path": "Freepik/F-Lite", "quantize": True, "quantize_te": True},
"sample": {"width": 1024, "height": 1024, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
"sample": {"width": 512, "height": 512, "num_inference_steps": 25, "guidance_scale": 4.0, "seed": 42},
},
}
@@ -228,7 +229,14 @@ def run_one(arch: str, device: str, allow_download: bool) -> dict:
from toolkit.config_modules import GenerateImageConfig, ModelConfig
from toolkit.util.get_model import get_model_class
model_config = ModelConfig(arch=arch, dtype="bf16", **entry["model"])
# quantization always tests the convrot8 backend (per-entry override wins)
model_kwargs = dict(entry["model"])
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)
ModelClass = get_model_class(model_config)
from toolkit.util.get_model import LEGACY_ARCHS
@@ -271,7 +279,32 @@ def run_one(arch: str, device: str, allow_download: bool) -> dict:
dtype="bf16",
noise_scheduler=sampler,
)
# ---- per-stage speed / VRAM / CPU instrumentation + cProfile ----
import cProfile
from testing.stage_profiler import StageProfiler, profile_top
prof = StageProfiler(device)
if hasattr(sd, "print_and_status_update"):
# every holder announces its stages through print_and_status_update;
# each status line opens a new profiler stage
_orig_status = sd.print_and_status_update
def _status_hook(msg, *a, **k):
prof.stage(str(msg))
return _orig_status(msg, *a, **k)
sd.print_and_status_update = _status_hook
prof.stage("load: init")
load_profile = cProfile.Profile()
load_profile.enable()
t_load0 = time.perf_counter()
sd.load_model()
load_seconds = time.perf_counter() - t_load0
load_profile.disable()
load_profile.dump_stats(os.path.join(out_dir, "load.prof"))
sample_kwargs = dict(entry["sample"])
if entry.get("needs_control_image"):
@@ -295,16 +328,175 @@ def run_one(arch: str, device: str, allow_download: bool) -> dict:
if not hasattr(ModelClass, "get_train_scheduler"):
# the legacy monolith takes the sampler NAME at generate time
gen_kwargs["sampler"] = "ddpm"
prof.stage("generate")
gen_profile = cProfile.Profile()
gen_profile.enable()
t_gen0 = time.perf_counter()
sd.generate_images([gen], **gen_kwargs)
gen_seconds = time.perf_counter() - t_gen0
gen_profile.disable()
gen_profile.dump_stats(os.path.join(out_dir, "gen.prof"))
stages = prof.finish()
produced = [
p
for p in glob.glob(os.path.join(out_dir, "*"))
if os.path.isfile(p) and os.path.getsize(p) > 1024 and not p.endswith(".txt")
if os.path.isfile(p)
and os.path.getsize(p) > 1024
and not p.endswith((".txt", ".prof", ".json"))
]
if not produced:
raise RuntimeError(f"no output file produced in {out_dir}")
return {"arch": arch, "status": "PASS", "outputs": produced}
import torch
result = {
"arch": arch,
"status": "PASS",
"outputs": produced,
"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),
"stages": stages,
"profile": {
"load_top": profile_top(load_profile),
"gen_top": profile_top(gen_profile, limit=15),
"load_prof": os.path.join(out_dir, "load.prof"),
"gen_prof": os.path.join(out_dir, "gen.prof"),
},
"env": {
"torch_num_threads": torch.get_num_threads(),
"cpu_count": os.cpu_count(),
"omp_num_threads": os.environ.get("OMP_NUM_THREADS"),
},
}
with open(os.path.join(out_dir, "metrics.json"), "w") as f:
json.dump(result, f, indent=2)
return result
def format_stage_table(stages) -> str:
"""Aligned markdown table: stage column left-aligned, numbers right-aligned
with fixed decimals so they line up in the terminal too."""
headers = [
"stage", "seconds", "vram peak GB", "vram resv GB", "vram end GB",
"rss peak GB", "cpu avg %", "cpu max %", "threads",
]
def fmt(value, decimals):
if value is None or value == "-":
return "-"
return f"{value:.{decimals}f}"
rows = [
[
s["name"],
fmt(s.get("seconds"), 2),
fmt(s.get("vram_peak_gb", "-"), 2),
fmt(s.get("vram_reserved_gb", "-"), 2),
fmt(s.get("vram_end_gb", "-"), 2),
fmt(s.get("rss_peak_gb"), 2),
fmt(s.get("cpu_avg_pct"), 1),
fmt(s.get("cpu_max_pct"), 1),
str(s.get("threads_max", "-")),
]
for s in stages
]
widths = [
max(len(headers[i]), *(len(r[i]) for r in rows)) if rows else len(headers[i])
for i in range(len(headers))
]
def line(cells):
out = []
for i, cell in enumerate(cells):
# stage name left-aligned, everything else right-aligned
out.append(cell.ljust(widths[i]) if i == 0 else cell.rjust(widths[i]))
return "| " + " | ".join(out) + " |"
sep = "|" + "|".join(
(":" + "-" * (w + 1)) if i == 0 else ("-" * (w + 1) + ":")
for i, w in enumerate(widths)
) + "|"
return "\n".join([line(headers), sep] + [line(r) for r in rows])
def write_report(results, path):
"""Aggregated markdown report: per-arch stage tables + cross-arch summary
+ load-profile hotspots, for diagnosing load/quantization speedups and
CPU usage."""
def _sum(stages, pred):
return round(sum(s["seconds"] for s in stages if pred(s)), 1)
lines = ["# Model loading test report", ""]
lines.append(
"Per-arch stage timings, peak VRAM (torch allocated/reserved), peak "
"process RSS, and CPU utilization (100 = one core, sampled at 50ms). "
"Quantization is convrot8 across the board. Full cProfile dumps sit "
"next to each arch's outputs (load.prof / gen.prof; inspect with "
"`python -m pstats <file>` or snakeviz)."
)
lines.append("")
# ---- summary ----
lines += ["## Summary", ""]
lines += [
"| arch | status | load s | quantize s | generate s | vram peak GB | rss peak GB | quant cpu avg % |",
"|---|---|---|---|---|---|---|---|",
]
for r in results:
if r["status"] != "PASS" or "stages" not in r:
lines.append(
f"| {r['arch']} | {r['status']} | - | - | - | - | - | - | "
)
continue
stages = r["stages"]
quant = [s for s in stages if "uantiz" in s["name"]]
vram_peak = max((s.get("vram_peak_gb", 0) for s in stages), default=0)
rss_peak = max((s.get("rss_peak_gb", 0) for s in stages), default=0)
quant_cpu = (
round(sum(s["cpu_avg_pct"] for s in quant) / len(quant), 1)
if quant
else "-"
)
lines.append(
f"| {r['arch']} | PASS | {r['load_seconds']} "
f"| {_sum(quant, lambda s: True)} | {r['gen_seconds']} "
f"| {vram_peak} | {rss_peak} | {quant_cpu} |"
)
lines.append("")
# ---- per-arch detail ----
for r in results:
if r["status"] != "PASS" or "stages" not in r:
lines += [f"## {r['arch']} — {r['status']}", "", r.get("error", ""), ""]
continue
lines += [f"## {r['arch']}", ""]
lines += [
f"load {r['load_seconds']}s, generate {r['gen_seconds']}s, "
f"qtype {r.get('qtype')}, qtype_te {r.get('qtype_te')}, "
f"torch threads {r['env']['torch_num_threads']}/{r['env']['cpu_count']} cores",
"",
]
lines += [format_stage_table(r["stages"]), ""]
top = r.get("profile", {}).get("load_top", [])[:12]
if top:
lines += ["Load hotspots (cumulative):", "", "```"]
for row in top:
lines.append(
f"{row['cumtime']:>9.2f}s tot {row['tottime']:>8.2f}s "
f"x{row['ncalls']:<8} {row['func']}"
)
lines += ["```", ""]
with open(path, "w") as f:
f.write("\n".join(lines))
return path
def main():
@@ -315,6 +507,11 @@ 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(
"--report-only",
action="store_true",
help="rebuild report.md from each arch's saved metrics.json, no runs",
)
args = parser.parse_args()
if args.list:
@@ -322,6 +519,24 @@ def main():
print(arch)
return
if args.report_only:
results = []
for arch in MODEL_TESTS:
out_dir = os.path.join(
OUTPUT_ROOT, arch.replace("/", "_").replace(":", "_")
)
metrics_path = os.path.join(out_dir, "metrics.json")
if os.path.exists(metrics_path):
with open(metrics_path) as f:
results.append(json.load(f))
else:
results.append(
{"arch": arch, "status": "SKIP", "error": "no metrics.json saved"}
)
report_path = write_report(results, os.path.join(OUTPUT_ROOT, "report.md"))
print(f"report: {report_path}")
return
os.environ.setdefault("CUDA_DEVICE_ORDER", "PCI_BUS_ID")
if not args.allow_download:
os.environ.setdefault("HF_HUB_OFFLINE", "1")
@@ -343,6 +558,10 @@ def main():
if args.json_result:
with open(args.json_result, "w") as f:
json.dump(result, f)
if result["status"] == "PASS" and "stages" in result:
print()
print(format_stage_table(result["stages"]))
print()
print(f"[{result['status']}] {args.arch}" + (f" — {result.get('error', '')}" if result["status"] != "PASS" else ""))
if result["status"] == "FAIL":
sys.exit(1)
@@ -355,8 +574,10 @@ def main():
# --all: one subprocess per arch so every model fully unloads (clean CUDA
# teardown) before the next loads
results = []
for arch in MODEL_TESTS:
print(f"\n===== {arch} =====")
total = len(MODEL_TESTS)
sweep_t0 = time.perf_counter()
for i, arch in enumerate(MODEL_TESTS, start=1):
print(f"\n===== {arch} ({i}/{total}) =====")
result_path = os.path.join(OUTPUT_ROOT, f".{arch.replace('/', '_')}.result.json")
cmd = [
sys.executable,
@@ -380,6 +601,18 @@ def main():
{"arch": arch, "status": "FAIL", "error": f"subprocess died (exit {proc.returncode})"}
)
# running tally so a long sweep always shows how far along it is
tally = {"PASS": 0, "FAIL": 0, "SKIP": 0}
for r in results:
tally[r["status"]] = tally.get(r["status"], 0) + 1
elapsed = time.perf_counter() - sweep_t0
eta = (elapsed / i) * (total - i)
print(
f">>> progress: {i}/{total} ({i * 100 // total}%) — "
f"{tally['PASS']} pass, {tally['FAIL']} fail, {tally['SKIP']} skip — "
f"elapsed {elapsed / 60:.0f}m, eta ~{eta / 60:.0f}m"
)
print("\n===== summary =====")
counts = {"PASS": 0, "FAIL": 0, "SKIP": 0}
for r in results:
@@ -389,6 +622,9 @@ def main():
line += f" — {r.get('error', '')[:160]}"
print(line)
print(f"\n{counts['PASS']} passed, {counts['FAIL']} failed, {counts['SKIP']} skipped")
report_path = write_report(results, os.path.join(OUTPUT_ROOT, "report.md"))
print(f"report: {report_path}")
if counts["FAIL"]:
sys.exit(1)

View File

@@ -91,6 +91,13 @@ class OstrisModelMixin:
# precision instead
aitk_cast_on_load: bool = True
# pre-quantized (comfy marker) checkpoints normally load their
# NON-quantized tensors at stored precision (the mix can be deliberate,
# e.g. ltx2.5's fp32 tables). Classes whose forward assumes one uniform
# dtype (wan: fp16 conv biases next to fp32 tables in the fp8 files) set
# this True to cast those float tensors to the load dtype instead
aitk_cast_quantized_load: bool = False
# ---- state set by the loader / quantizer ----
aitk_is_quantized: bool = False
aitk_qtype: Optional[str] = None
@@ -461,7 +468,22 @@ class OstrisModelMixin:
state_dict, num_quantized = import_comfy_quantized_layers(
model, state_dict, orig_dtype=dtype
)
if cls.aitk_cast_quantized_load:
# uniform-dtype models: the leftover (non-quantized) float
# tensors follow the load dtype, like the non-marker path
for key, value in state_dict.items():
if value.is_floating_point():
state_dict[key] = value.to(dtype=dtype)
cls._load_state_dict_with_quantized(model, state_dict)
if cls.aitk_cast_quantized_load:
# the importer assigns quantized-layer biases directly at
# stored dtype; they must follow too (a stray fp16 bias also
# makes ModelMixin.dtype — the pipelines' cast target — lie)
from toolkit.util.ostris_quant import OstrisLinear
for m in model.modules():
if isinstance(m, OstrisLinear) and m.bias is not None:
m.bias.data = m.bias.data.to(dtype=dtype)
model.aitk_is_quantized = True
elif cls.aitk_cast_on_load:
for key, value in state_dict.items():

View File

@@ -76,6 +76,12 @@ class WanTransformer3DModel(DiffusersWanTransformer3DModel, OstrisModelMixin):
},
}
# the comfy fp8/scaled_fp8 wan files mix fp16 conv biases with fp32
# tables; wan's forward assumes one uniform dtype, so cast the
# non-quantized tensors to the load dtype (the old holder did this with a
# blanket .to(device, dtype) after load)
aitk_cast_quantized_load = True
@classmethod
def get_transformer_block_names(cls):
return ["blocks"]

View File

@@ -34,6 +34,9 @@ class UMT5TextEncoder(UMT5EncoderModel, OstrisTransformersMixin):
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
aitk_config_repo = "ai-toolkit/umt5_xxl_encoder"
# comfy fp8 umt5 files store non-quantized tensors in a dtype mix; T5
# assumes a uniform dtype, so follow the load dtype
aitk_cast_quantized_load = True
aitk_comfy_repo = "Comfy-Org/Wan_2.1_ComfyUI_repackaged"
# comfy umt5 files already use the transformers key layout; the

View File

@@ -1319,6 +1319,19 @@ class ConvRotInt8Quantizer(OstrisQuantizer):
f"(needs in divisible by 16, out by 8, and a power-of-4 block >= 16)"
)
return False
if d < 128 and module.out_features % 16 != 0:
# cublasLt has no int8 kernel for K < 128 with N % 16 != 0
# (torch._int_mm raises CUBLAS_STATUS_NOT_SUPPORTED), e.g.
# omnigen2's 64 -> 2520 x_embedder
key = (d, module.out_features)
if key not in _skip_warned:
_skip_warned.add(key)
print_acc(
f"ConvRot: skipping linears with in_features={d}, "
f"out_features={module.out_features} (int8 gemm needs out "
f"divisible by 16 when in < 128)"
)
return False
return True
def quantize_(self, module: torch.nn.Linear, weight_fp32: torch.Tensor) -> None: