Fix quantization issues
This commit is contained in:
@@ -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
157
testing/stage_profiler.py
Normal 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
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user