Optimize quantization path
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user