diff --git a/extensions_built_in/diffusion_models/hidream/hidream_o1_model.py b/extensions_built_in/diffusion_models/hidream/hidream_o1_model.py index 41bcf0c..1011808 100644 --- a/extensions_built_in/diffusion_models/hidream/hidream_o1_model.py +++ b/extensions_built_in/diffusion_models/hidream/hidream_o1_model.py @@ -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, ) diff --git a/testing/stage_profiler.py b/testing/stage_profiler.py new file mode 100644 index 0000000..ab12020 --- /dev/null +++ b/testing/stage_profiler.py @@ -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 ("",) 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 == "": + 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 diff --git a/testing/test_model_loading.py b/testing/test_model_loading.py index 56da4f9..a1f45b1 100644 --- a/testing/test_model_loading.py +++ b/testing/test_model_loading.py @@ -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 ` 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) diff --git a/toolkit/models/v2/_mixin.py b/toolkit/models/v2/_mixin.py index def4cbb..5a2e71e 100644 --- a/toolkit/models/v2/_mixin.py +++ b/toolkit/models/v2/_mixin.py @@ -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(): diff --git a/toolkit/models/v2/diffusion_models/wan.py b/toolkit/models/v2/diffusion_models/wan.py index d71cdaf..fd0367a 100644 --- a/toolkit/models/v2/diffusion_models/wan.py +++ b/toolkit/models/v2/diffusion_models/wan.py @@ -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"] diff --git a/toolkit/models/v2/text_encoders/umt5.py b/toolkit/models/v2/text_encoders/umt5.py index 8bd02e0..2b86dff 100644 --- a/toolkit/models/v2/text_encoders/umt5.py +++ b/toolkit/models/v2/text_encoders/umt5.py @@ -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 diff --git a/toolkit/util/convrot_quant.py b/toolkit/util/convrot_quant.py index 0411281..ec0645a 100644 --- a/toolkit/util/convrot_quant.py +++ b/toolkit/util/convrot_quant.py @@ -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: