Optimize quantization path

This commit is contained in:
Jaret Burkett
2026-08-28 11:47:29 -06:00
parent ce3df32101
commit 5ddc5f8ca7
3 changed files with 127 additions and 9 deletions

View File

@@ -442,6 +442,40 @@ them: Flux2, MageFlow, ZImageDCT, Ideogram4Transformer2DModel (+ its exclude
list on MageFlow). `prepare_text_encoder` remains only as the legacy-monolith
path.
## Load/quantize optimization round 1 (2026-08-28)
Baseline: testing/.model_test_outputs/report_baseline.md (+ per-arch
metrics_baseline.json). Profiling findings and fixes:
- The "slow quantization" in the baseline was mostly NOT quantize math
(convrot8 on chroma's 57 blocks: 0.5s of GPU kernels). It was data
movement: mmap-backed safetensors faulting pages in per-tensor cudaMemcpy
order (cold cache on the 92%-full NVMe reads at ~0.32GB/s; the drive's
sequential ceiling is ~435MB/s — cold loads are storage-bound), plus the
old stream-quantize choreography crossing the bus three times (up, back
into a fresh host copy, up again at final placement).
- `quantize_module(keep_on_device=True)`: when the final home IS the
quantize gpu (no offload/low_vram — aitk_post_load detects this), blocks
stay on the gpu after quantizing and the extras pass runs there too.
Weights cross the bus once; the model-sized host-RAM spike of the
intermediate copy is gone (chroma RSS peak 25.3 -> 17.3GB, the rest is
reclaimable page cache; flux2's 93GB spike class eliminated). Warm-cache
chroma blocks: 46s -> 3.9s; wan22_5b quantize 10.6s -> 3.8s.
- Offload/low_vram path keeps cpu residency but quantizes extras layer-by-
layer via quantize_device instead of the all-cores cpu burst (was 350-650%
cpu avg with 1600% spikes).
- `quantize()` skips the quantize_device round-trip for same-qtype
OstrisLinears (guaranteed no-op) and, for ostris backends, for non-Linear
leaves (norms/embeddings were being ferried across the bus for nothing).
- posix_fadvise(WILLNEED) readahead on single-file loads (mixin._readahead)
— ~10-15% on cold reads, free when warm, no-op off POSIX.
Ideas not yet done: pinned-buffer staged h2d (~+60% warm-upload bandwidth,
9.5 vs 5.8GB/s measured), overlapping h2d with quantize kernels (small — the
kernels are ~2% of the pass), wan22 low_vram generate choreography (165s of
unload/reload per sample), extending keep_on_device to the legacy
quantize_model/trainer path.
## Testing
- [x] `testing/test_model_loading.py`: per-arch load + one small sample through

View File

@@ -257,10 +257,21 @@ class OstrisModelMixin:
attach_ara_and_quantize(base_model, self, ara_path, exclude=exclude)
else:
status_fn(f"Quantizing ({qtype})")
q_device = quantize_device if quantize_device is not None else device
# final home is the quantize gpu (no offloading, no cpu
# parking): quantized weights stay put — one bus crossing
# instead of three, and no model-sized host-ram copy
keep_on_device = (
not offload
and device is not None
and q_device is not None
and torch.device(device) == torch.device(q_device)
and torch.device(device).type != "cpu"
)
quantize_module(
self,
qtype,
device=quantize_device if quantize_device is not None else device,
device=q_device,
dtype=dtype,
block_names=self.get_transformer_block_names(),
exclude=exclude,
@@ -270,6 +281,7 @@ class OstrisModelMixin:
else None
),
status_fn=status_fn,
keep_on_device=keep_on_device,
)
self.aitk_is_quantized = True
self.aitk_qtype = qtype
@@ -407,6 +419,22 @@ class OstrisModelMixin:
subfolder = None
return cls.aitk_load_config(config_source, subfolder=subfolder)
@staticmethod
def _readahead(file_path: str):
"""Queue kernel readahead for the whole file. safetensors tensors are
mmap-backed; without this a cold cache faults pages in per-tensor
cudaMemcpy order at a fraction of the drive's sequential bandwidth."""
if not hasattr(os, "posix_fadvise"):
return # non-POSIX: mmap readahead heuristics only
try:
fd = os.open(file_path, os.O_RDONLY)
try:
os.posix_fadvise(fd, 0, 0, os.POSIX_FADV_WILLNEED)
finally:
os.close(fd)
except OSError:
pass
@classmethod
def _load_single_file(
cls,
@@ -417,6 +445,7 @@ class OstrisModelMixin:
subfolder: Optional[str] = None,
**kwargs,
):
cls._readahead(file_path)
state_dict = load_file(file_path)
return cls.load_from_state_dict(
state_dict,

View File

@@ -1,3 +1,5 @@
import os
import time
from fnmatch import fnmatch
from typing import List, Optional, Union, TYPE_CHECKING
import torch
@@ -186,6 +188,16 @@ def quantize(
and m.__class__.__name__ == "OstrisLinear"
):
continue
if getattr(m.ostris_quantizer, "qtype", None) == weights.quantizer.qtype:
# guaranteed per-layer no-op: skip before any
# quantize_device round-trip
continue
if isinstance(weights, ostristype) and not isinstance(m, torch.nn.Linear):
# ostris backends only quantize nn.Linear; don't ferry norms/
# embeddings across the bus for nothing when quantize_device
# is set (containers fall through so children are visited)
if quantize_device is not None and next(m.children(), None) is None:
continue
if (
isinstance(weights, aotype)
and not isinstance(m, torch.nn.Linear)
@@ -271,17 +283,29 @@ def quantize_module(
exclude: Optional[List[str]] = None,
quantize_kwargs: Optional[dict] = None,
status_fn=print_acc,
keep_on_device: bool = False,
):
"""Module-centric quantization: block-streamed (each repeated block moves
to ``device`` for the math and returns to cpu) with a whole-module pass
for the extras. This is the core the per-model loaders call; holders'
quantize_model wraps it with model_config plumbing."""
to ``device`` for the math) with a whole-module pass for the extras. This
is the core the per-model loaders call; holders' quantize_model wraps it
with model_config plumbing.
``keep_on_device``: when the model's final home IS ``device`` (no layer
offloading / low_vram), quantized blocks stay there instead of round-
tripping back to cpu — the weights then cross the bus exactly once
(mmap/page-cache -> gpu) instead of three times (up, back into a fresh
host copy, up again at final placement), and the model-sized host RAM
spike of that intermediate copy never happens. The extras pass runs on
the gpu too. With it off (offload paths), blocks return to cpu as before
and the extras quantize layer-by-layer on ``device`` via quantize_device
rather than burning every cpu core."""
from toolkit.dequantize import patch_dequantization_on_save
patch_dequantization_on_save(module)
quantization_type = get_qtype(qtype)
exclude = list(exclude or [])
quantize_kwargs = quantize_kwargs or {}
keep_on_device = keep_on_device and device is not None
all_blocks: List[torch.nn.Module] = []
for name in block_names or []:
@@ -296,29 +320,60 @@ def quantize_module(
if all_blocks:
status_fn(f" - quantizing {len(all_blocks)} blocks")
already_quantized = 0
debug_phases = os.environ.get("AITK_QUANT_DEBUG") == "1"
t_check = t_h2d = t_quant = t_d2h = 0.0
for block in tqdm(all_blocks):
if not _has_quantizable_linear(block, quantization_type, exclude):
t = time.perf_counter()
skip = not _has_quantizable_linear(block, quantization_type, exclude)
t_check += time.perf_counter() - t
if skip:
# pre-quantized checkpoint with a matching qtype: nothing in this
# block would change — skip the device round-trip and the dtype
# cast entirely so the load stays byte-identical
# block would change — skip the dtype cast entirely so the load
# stays byte-identical (placement still honors keep_on_device)
already_quantized += 1
if keep_on_device:
block.to(device)
continue
t = time.perf_counter()
if device is not None:
block.to(device, dtype=dtype, non_blocking=True)
t_h2d += time.perf_counter() - t
t = time.perf_counter()
quantize(block, weights=quantization_type, exclude=exclude, **quantize_kwargs)
freeze(block)
t_quant += time.perf_counter() - t
# NOT non_blocking: an async D2H allocates the cpu destination in pinned
# memory, which the caching host allocator keeps forever — that silently
# retained a model-sized chunk of host ram
if device is not None:
t = time.perf_counter()
if device is not None and not keep_on_device:
block.to("cpu")
t_d2h += time.perf_counter() - t
if debug_phases and all_blocks:
status_fn(
f" - [debug] check {t_check:.1f}s h2d {t_h2d:.1f}s "
f"quant {t_quant:.1f}s d2h {t_d2h:.1f}s"
)
if already_quantized:
status_fn(
f" - {already_quantized} blocks already quantized with a matching qtype; left untouched"
)
status_fn(" - quantizing extras")
quantize(module, weights=quantization_type, exclude=exclude, **quantize_kwargs)
if keep_on_device:
# blocks already live on the gpu; bring the remainder over and run the
# extras pass there (fast kernels instead of an all-cores cpu burst)
module.to(device)
quantize(module, weights=quantization_type, exclude=exclude, **quantize_kwargs)
else:
# cpu-resident model: quantize each extra layer with a gpu round-trip
quantize(
module,
weights=quantization_type,
exclude=exclude,
quantize_device=device,
**quantize_kwargs,
)
freeze(module)
return module