From bc1d2143ca16319a9e6f5048b5cf37ef999ceaba Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Sat, 18 Jul 2026 16:24:46 +0200 Subject: [PATCH 1/8] Quiet non-actionable startup warnings Suppress a handful of specific, noisy-but-harmless messages emitted while launching the UI and starting training: - diffusers/transformers logger.warning() lines (Modular Diffusers experimental notice, unexpected-config-attributes, unrecognized loss_type) via filters on the exact emitting loggers - huggingface_hub local_dir_use_symlinks deprecation and the torch.compile inductor performance notes via warnings/logger filters - Qt gnome portal dbus errors via QT_LOGGING_RULES - tensorboard subprocess banner/notices by discarding its stdout/stderr Each filter targets one specific message, so other warnings from the same libraries still come through. Co-Authored-By: Claude Opus 4.8 (1M context) --- modules/trainer/BaseTrainer.py | 8 +++++++- modules/ui/TrainUIController.py | 7 ++++++- modules/util/ui/pyside6_util.py | 9 +++++++++ scripts/util/import_util.py | 36 +++++++++++++++++++++++++++++++++ 4 files changed, 58 insertions(+), 2 deletions(-) diff --git a/modules/trainer/BaseTrainer.py b/modules/trainer/BaseTrainer.py index 20fc5eb7a..e4b7a3b29 100644 --- a/modules/trainer/BaseTrainer.py +++ b/modules/trainer/BaseTrainer.py @@ -97,7 +97,13 @@ def _start_tensorboard(self): if self.config.tensorboard_expose: tensorboard_args.append("--bind_all") - self.tensorboard_subprocess = subprocess.Popen(tensorboard_args) + # Discard the tensorboard child's stdout/stderr: the TF-not-found notice, the + # experimental-data-loading NOTE and the serving banner are all noise, and the + # UI already exposes the tensorboard URL. Popen still raises if the executable + # is missing, so a real launch failure is not hidden. + self.tensorboard_subprocess = subprocess.Popen( + tensorboard_args, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, + ) def _stop_tensorboard(self): self.tensorboard_subprocess.kill() diff --git a/modules/ui/TrainUIController.py b/modules/ui/TrainUIController.py index 72623def4..2c4896226 100644 --- a/modules/ui/TrainUIController.py +++ b/modules/ui/TrainUIController.py @@ -103,8 +103,13 @@ def _start_always_on_tensorboard(self): if self.train_config.tensorboard_expose: tensorboard_args.append("--bind_all") + # Discard the tensorboard child's stdout/stderr: the TF-not-found notice, the + # experimental-data-loading NOTE and the serving banner are all noise, and the + # UI already exposes the tensorboard URL. try: - self.always_on_tensorboard_subprocess = subprocess.Popen(tensorboard_args) + self.always_on_tensorboard_subprocess = subprocess.Popen( + tensorboard_args, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, + ) except Exception: self.always_on_tensorboard_subprocess = None diff --git a/modules/util/ui/pyside6_util.py b/modules/util/ui/pyside6_util.py index fd7b2ef8a..42b4a3e60 100644 --- a/modules/util/ui/pyside6_util.py +++ b/modules/util/ui/pyside6_util.py @@ -1,3 +1,4 @@ +import os import signal import sys from abc import ABCMeta @@ -19,6 +20,14 @@ def create_application() -> QApplication: # active and Ctrl+C would be ignored. signal.signal(signal.SIGINT, signal.SIG_DFL) + # On desktops without the xdg-desktop-portal Settings interface, Qt spams two + # "qt.qpa.theme.gnome: dbus reply error ... org.freedesktop.portal.Settings" + # lines while probing for the system color scheme. Silence just that category; + # the rules string is read when Qt's logging initializes at QApplication init. + _gnome_theme_rule = "qt.qpa.theme.gnome=false" + existing_rules = os.environ.get("QT_LOGGING_RULES") + os.environ["QT_LOGGING_RULES"] = f"{existing_rules};{_gnome_theme_rule}" if existing_rules else _gnome_theme_rule + app = QApplication(sys.argv) # Force Fusion everywhere: native styles (e.g. windowsvista) draw standard # controls via OS theme APIs, which breaks once an application stylesheet diff --git a/scripts/util/import_util.py b/scripts/util/import_util.py index 150df4d8a..a0c0a83b6 100644 --- a/scripts/util/import_util.py +++ b/scripts/util/import_util.py @@ -1,7 +1,9 @@ def script_imports(allow_zluda: bool = True): import logging import os + import re import sys + import warnings from pathlib import Path # Filter out the Triton warning on startup. @@ -10,6 +12,40 @@ def script_imports(allow_zluda: bool = True): .getLogger("xformers") \ .addFilter(lambda record: 'A matching Triton is not available' not in record.getMessage()) + # Silence specific non-actionable startup/compile warnings. A logger filter + # targets the exact emitting logger, since a parent logger's filter misses + # records from child loggers. + + # diffusers/transformers chatty logger.warning() lines at import/load time. + logging.getLogger("diffusers.modular_pipelines").addFilter( + lambda record: 'Modular Diffusers is currently an experimental feature' not in record.getMessage() + ) + # The subject of these two is interpolated into the message, so match the whole + # sentence with .* standing in for the runtime value. + logging.getLogger("diffusers.configuration_utils").addFilter( + lambda record: not re.search( + r"The config attributes .* were passed to .*, but are not expected and will be ignored", + record.getMessage(), + ) + ) + logging.getLogger("transformers.modeling_utils").addFilter( + lambda record: not re.search( + r"`loss_type=.*` was set in the config but it is unrecognized", record.getMessage() + ) + ) + + # A dependency still calls hf_hub_download with the removed local_dir_use_symlinks + # argument; the deprecation warning is not actionable. + warnings.filterwarnings("ignore", message=r".*local_dir_use_symlinks.*") + + # torch.compile emits performance notes when inductor falls back or can't use a + # fast path; harmless and noisy for normal runs. The SMs note is a logger.warning() + # on its exact emitting logger; the complex-operators note is a warnings.warn(). + warnings.filterwarnings("ignore", message=r".*does not support code generation for complex operators.*") + logging.getLogger("torch._inductor.utils").addFilter( + lambda record: 'Not enough SMs to use max_autotune_gemm mode' not in record.getMessage() + ) + # Insert ourselves as the highest-priority library path, so our modules are # always found without any risk of being shadowed by another import path. # 3 .parent calls to navigate from /scripts/util/import_util.py to the main directory From b9ac6afcad6d0979fe03fc9221711fe956b653dc Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Sat, 8 Aug 2026 12:46:00 +0200 Subject: [PATCH 2/8] W8A8 kernel optimizations and LoRA fusion Reworks the Triton 8-bit matmuls around a shared compute core and adds a LoRA epilogue to it, so a LoRA layer over an 8-bit quantized Linear no longer materializes its low-rank product separately. The three mm entry points now share _mm_accumulate (grouped launch order for L2 reuse, compile-time layout specialization, an EVEN_K-specialized main loop) and differ only in their epilogue: the raw product, a per-row dequant scale, or that scale plus a rank-tiled low-rank update. The scaled_lora_mm_8bit and rowcol_scaled_lora_mm_8bit entry points expose the last one. transpose_8bit rewrites the backward pass's B operand to k-major, since the 8-bit tensor-core op wants k-major and Ada has no 8-bit ldmatrix.trans; the wrappers apply it above a token threshold, where the copy pays for itself. On the module side, LoRAFusableLinearMixin declares forward_with_lora, and LinearW8A8, LinearGGUFA8 and LinearSVD implement it with autograd Functions that own the down-projection and the dropout. That lets the LoRA dgrad fold into the backward epilogue as well, so neither direction builds an (M, out_features) intermediate. LoRAModule dispatches on the mixin rather than on BaseLinearSVD, and fused_leaf_forward extends the same path to fused-qkv adapters by narrowing lora_up to the leaf's rows. DoRA opts out, as it recomposes the weight instead of adding a delta. mm_8bit falls back to torch._int_mm / torch._scaled_mm when Triton is not importable, and quantize_axiswise dispatches the two 8-bit dtypes. Co-Authored-By: Claude Opus 5 (1M context) --- modules/module/FusedModule.py | 5 + modules/module/LoRAModule.py | 25 +- modules/module/quantized/LinearGGUFA8.py | 174 ++++-- modules/module/quantized/LinearSVD.py | 43 +- modules/module/quantized/LinearW8A8.py | 235 ++++++-- .../quantized/mixin/LoRAFusableLinearMixin.py | 9 + modules/util/mm_8bit.py | 26 +- modules/util/quantization_util.py | 8 + modules/util/triton_mm_8bit.py | 513 +++++++++++++++--- 9 files changed, 871 insertions(+), 167 deletions(-) create mode 100644 modules/module/quantized/mixin/LoRAFusableLinearMixin.py diff --git a/modules/module/FusedModule.py b/modules/module/FusedModule.py index 1948c7270..f517c4102 100644 --- a/modules/module/FusedModule.py +++ b/modules/module/FusedModule.py @@ -95,6 +95,11 @@ def __init__(self, prefix: str, leaves: list[nn.Module], klass, additional_args: # recompose the base weight itself (delta_forward returns None) keep going through the slower, # generic self.module.forward(x) path. def _leaf_output(self, leaf_index: int, start: int, end: int, x, *args, **kwargs): + # a leaf that can fold the adapter into its own base matmul returns the finished output, so + # neither the full fused delta nor the separate add happens for it + fused = self.module.fused_leaf_forward(self.leaves[leaf_index], x, start, end) + if fused is not None: + return fused delta = self.module.delta_forward(x, *args, **kwargs) if delta is None: return self.module.forward(x)[..., start:end] diff --git a/modules/module/LoRAModule.py b/modules/module/LoRAModule.py index 4f20c56e4..81370b864 100644 --- a/modules/module/LoRAModule.py +++ b/modules/module/LoRAModule.py @@ -7,7 +7,7 @@ from modules.module.FusedModule import FusedModuleGroup, check_fusion_match, discover_fused_groups from modules.module.oft_utils import OFTRotationModule -from modules.module.quantized.LinearSVD import BaseLinearSVD +from modules.module.quantized.mixin.LoRAFusableLinearMixin import LoRAFusableLinearMixin from modules.util.config.TrainConfig import TrainConfig from modules.util.enum.ModelType import PeftType from modules.util.lokr_utils import factorization, make_kron, rebuild_tucker @@ -115,6 +115,13 @@ def delta_forward(self, x, *args, **kwargs) -> Tensor | None: # going through the fused forward. return None + def fused_leaf_forward(self, leaf: nn.Module, x, start: int, end: int) -> Tensor | None: + # Returns one leaf's complete output (base + this adapter's contribution) when the leaf can + # fold the adapter into its own base matmul, so no delta is computed or added separately. + # start/end are the leaf's row range in the fused output. Returns None (the default) when no + # such path exists, and the caller falls back to delta_forward. + return None + @property def orig_module(self) -> nn.Module: assert self._orig_module is not None @@ -570,11 +577,18 @@ def check_initialized(self): def forward(self, x, *args, **kwargs): self.check_initialized() - if isinstance(self.orig_module, BaseLinearSVD): - return self.orig_module.forward_with_lora(x, self.lora_down, self.lora_up, self.dropout, self.alpha) + if isinstance(self.orig_module, LoRAFusableLinearMixin): + return self.orig_module.forward_with_lora(x, self.lora_down.weight, self.lora_up.weight, self.dropout, self.alpha) return self.orig_forward(x) + self.delta_forward(x, *args, **kwargs) + def fused_leaf_forward(self, leaf: nn.Module, x, start: int, end: int) -> Tensor | None: + self.check_initialized() + if not isinstance(leaf, LoRAFusableLinearMixin): + return None + #the fused adapter's up spans all leaves' concatenated outputs; narrow it to this leaf's rows + return leaf.forward_with_lora(x, self.lora_down.weight, self.lora_up.weight[start:end], self.dropout, self.alpha) + def delta_forward(self, x, *args, **kwargs) -> Tensor | None: self.check_initialized() ld = self.lora_up(self.dropout(self.lora_down(x))) @@ -781,6 +795,11 @@ def delta_forward(self, x, *args, **kwargs) -> Tensor | None: # DoRA scales the recomposed weight, so there is no delta term; back to None from LoRAModule's. return None + def fused_leaf_forward(self, leaf: nn.Module, x, start: int, end: int) -> Tensor | None: + # same reason as delta_forward: the leaf folds an additive term into its own base matmul, + # and DoRA recomposes the weight instead + return None + def forward(self, x, *args, **kwargs): self.check_initialized() A = self.lora_down.weight diff --git a/modules/module/quantized/LinearGGUFA8.py b/modules/module/quantized/LinearGGUFA8.py index 8b015c564..d11606d53 100644 --- a/modules/module/quantized/LinearGGUFA8.py +++ b/modules/module/quantized/LinearGGUFA8.py @@ -1,5 +1,8 @@ +from modules.module.quantized.mixin.LoRAFusableLinearMixin import LoRAFusableLinearMixin from modules.util.mm_8bit import mm_8bit as mm_8bit +from modules.util.mm_8bit import rowcol_scaled_lora_mm_8bit, rowcol_scaled_mm_8bit from modules.util.quantization_util import ( + quantize_axiswise, quantize_fp8_axiswise, quantize_int8_axiswise, ) @@ -14,70 +17,133 @@ UNQUANTIZED_TYPES = [gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16, gguf.GGMLQuantizationType.BF16] +#unlike LinearW8A8, whose weight is quantized once and tensorwise, the GGUF weight is dequantized and +#requantized axiswise per pass (it can be, since the dequant happens anyway). That leaves two scale +#vectors varying along different axes - per-token on M and per-output-channel on N - which cannot be +#collapsed into one, so this layer uses the rowcol_ variants of the epilogue-scaled kernels. + +#the weight reaches these dequantized, so unlike LinearW8A8 they are told which 8-bit dtype to +#requantize to. Only the torch forward differs beyond the quantizer: int8 and fp8 have separate +#torch mms with separate scaling @torch.no_grad() -def int8_forward_axiswise(x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: - x_8, x_scale = quantize_int8_axiswise(x, dim=-1) - w_8, w_scale = quantize_int8_axiswise(weight, dim=-1) - res = torch._int_mm(x_8, w_8.T) - res_scaled = res.float().mul_(w_scale.T).mul_(x_scale).to(compute_dtype) +def forward_axiswise_postscaled_torch(dtype: torch.dtype, x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: + if dtype == torch.int8: + x_8, x_scale = quantize_int8_axiswise(x, dim=-1) + w_8, w_scale = quantize_int8_axiswise(weight, dim=-1) + res = torch._int_mm(x_8, w_8.T) + res_scaled = res.float().mul_(w_scale.T).mul_(x_scale).to(compute_dtype) + else: + x_8, x_scale = quantize_fp8_axiswise(x, dim=-1) + w_8, w_scale = quantize_fp8_axiswise(weight, dim=-1) + one = torch.ones(1, device=x.device) + res = torch._scaled_mm(x_8, w_8.T, scale_a=one, scale_b=one, out_dtype=torch.float) + res_scaled = res.mul_(w_scale.T).mul_(x_scale).to(compute_dtype) #much faster than scaled by _scaled_mm if bias is not None: res_scaled.add_(bias) return res_scaled @torch.no_grad() -def fp8_forward_axiswise(x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: - x_8, x_scale = quantize_fp8_axiswise(x, dim=-1) - w_8, w_scale = quantize_fp8_axiswise(weight, dim=-1) - one = torch.ones(1, device=x.device) - res = torch._scaled_mm(x_8, w_8.T, scale_a=one, scale_b=one, out_dtype=torch.float) - res_scaled = res.mul_(w_scale.T).mul_(x_scale).to(compute_dtype) #much faster than scaled by _scaled_mm +def backward_axiswise_postscaled_triton(dtype: torch.dtype, output: Tensor, weight: Tensor) -> Tensor: + output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=dtype) + w_8, w_scale = quantize_axiswise(weight, dim=0, dtype=dtype) + mm_res = mm_8bit(output_8.contiguous(), w_8) + return mm_res.float().mul_(w_scale).mul_(output_scale).to(output.dtype) + + +@torch.no_grad() +def forward_axiswise_epiloguescaled_triton(dtype: torch.dtype, x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: + x_8, x_scale = quantize_axiswise(x, dim=-1, dtype=dtype) + w_8, w_scale = quantize_axiswise(weight, dim=-1, dtype=dtype) + #rowcol_scaled_mm_8bit folds the per-token scale (axis 0) and the per-channel weight scale + #(axis 1) into the epilogue and returns compute_dtype directly + res_scaled = rowcol_scaled_mm_8bit(x_8, w_8.T, x_scale, w_scale, compute_dtype) if bias is not None: res_scaled.add_(bias) return res_scaled @torch.no_grad() -def int8_backward_axiswise(output: Tensor, weight: Tensor) -> Tensor: - output_8, output_scale = quantize_int8_axiswise(output, dim=-1) - w_8, w_scale = quantize_int8_axiswise(weight, dim=0) - mm_res = mm_8bit(output_8.contiguous(), w_8) - return mm_res.float().mul_(w_scale).mul_(output_scale).to(output.dtype) +def backward_axiswise_epiloguescaled_triton(dtype: torch.dtype, output: Tensor, weight: Tensor) -> Tensor: + output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=dtype) + w_8, w_scale = quantize_axiswise(weight, dim=0, dtype=dtype) + return rowcol_scaled_mm_8bit(output_8.contiguous(), w_8, output_scale, w_scale, output.dtype) + @torch.no_grad() -def fp8_backward_axiswise(output: Tensor, weight: Tensor) -> Tensor: - output_8, output_scale = quantize_fp8_axiswise(output, dim=-1) - w_8, w_scale = quantize_fp8_axiswise(weight, dim=0) - mm_res = mm_8bit(output_8.contiguous(), w_8) - return mm_res.float().mul_(w_scale).mul_(output_scale).to(output.dtype) +def forward_axiswise_lora_epiloguescaled_triton(dtype: torch.dtype, x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype, x_down: Tensor, lora_up: Tensor) -> Tensor: + x_8, x_scale = quantize_axiswise(x, dim=-1, dtype=dtype) + w_8, w_scale = quantize_axiswise(weight, dim=-1, dtype=dtype) + res_scaled = rowcol_scaled_lora_mm_8bit(x_8, w_8.T, x_scale, w_scale, x_down, lora_up, compute_dtype) + if bias is not None: + res_scaled.add_(bias) + return res_scaled + +@torch.no_grad() +def backward_axiswise_lora_epiloguescaled_triton(dtype: torch.dtype, output: Tensor, weight: Tensor, grad_x_down_pre: Tensor, lora_down: Tensor) -> Tensor: + output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=dtype) + w_8, w_scale = quantize_axiswise(weight, dim=0, dtype=dtype) + return rowcol_scaled_lora_mm_8bit(output_8.contiguous(), w_8, output_scale, w_scale, grad_x_down_pre, lora_down.to(grad_x_down_pre.dtype), output.dtype) + + +forward_axiswise = forward_axiswise_epiloguescaled_triton +backward_axiswise = backward_axiswise_epiloguescaled_triton +forward_axiswise_lora = forward_axiswise_lora_epiloguescaled_triton +backward_axiswise_lora = backward_axiswise_lora_epiloguescaled_triton -class LinearGGUFIntA8RequantFunction(torch.autograd.Function): + +class LinearGGUFA8RequantFunction(torch.autograd.Function): + #`weight` is the dequantized GGUF weight, so saving it lets backward requantize instead of + #dequantizing the GGUF blocks a second time @staticmethod - def forward(ctx, x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: + def forward(ctx, x: Tensor, weight: Tensor, dtype: torch.dtype, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: ctx.save_for_backward(weight) - #axiswise performs better than tensorwise in tests, even though - #it requires another requant during backward - but requant is cheap - return int8_forward_axiswise(x, weight, bias, compute_dtype) + ctx.dtype = dtype + return forward_axiswise(dtype, x, weight, bias, compute_dtype) @staticmethod def backward(ctx, output: Tensor): - if ctx.needs_input_grad != (True, False, False, False): + if ctx.needs_input_grad != (True, False, False, False, False): raise NotImplementedError("GGUF cannot be used for full finetuning") + weight, = ctx.saved_tensors - return int8_backward_axiswise(output, weight), None, None, None + return backward_axiswise(ctx.dtype, output, weight), None, None, None, None + -class LinearGGUFFpA8RequantFunction(torch.autograd.Function): +class LinearGGUFA8RequantLoRAFunction(torch.autograd.Function): + #LinearGGUFA8RequantFunction plus a low-rank update: it owns the down-projection and the + #dropout, so the LoRA dgrad folds into the backward epilogue and the (M, out) product is + #never materialized @staticmethod - def forward(ctx, x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: - ctx.save_for_backward(weight) - return fp8_forward_axiswise(x, weight, bias, compute_dtype) + def forward(ctx, x: Tensor, weight: Tensor, dtype: torch.dtype, bias: Tensor | None, compute_dtype: torch.dtype, lora_down: Tensor, lora_up: Tensor, dropout_mask: Tensor | None) -> Tensor: + x_down = torch.nn.functional.linear(x, lora_down) + if dropout_mask is not None: + x_down = x_down * dropout_mask + ctx.dtype = dtype + ctx.save_for_backward(weight, x.to(x_down.dtype), lora_down, lora_up, dropout_mask, x_down) + return forward_axiswise_lora(dtype, x, weight, bias, compute_dtype, x_down, lora_up) @staticmethod - def backward(ctx, output: Tensor): - if ctx.needs_input_grad != (True, False, False, False): + def backward(ctx, grad_output: Tensor): + if ctx.needs_input_grad[1:5] != (False, False, False, False): raise NotImplementedError("GGUF cannot be used for full finetuning") - weight, = ctx.saved_tensors - return fp8_backward_axiswise(output, weight), None, None, None -class LinearGGUFA8(GGUFLinear): + weight, x, lora_down, lora_up, dropout_mask, x_down = ctx.saved_tensors + needs_x, needs_down, needs_up = ctx.needs_input_grad[0], ctx.needs_input_grad[5], ctx.needs_input_grad[6] + #grad_x_down_pre (post-dropout-backward) feeds both the grad_x epilogue fold and grad_lora_down + grad_x_down_pre = None + if needs_x or needs_down: + grad_x_down_pre = grad_output @ lora_up.T + if dropout_mask is not None: + grad_x_down_pre = grad_x_down_pre * dropout_mask + grad_x = backward_axiswise_lora(ctx.dtype, grad_output, weight, grad_x_down_pre, lora_down) if needs_x else None + grad_lora_down = (grad_x_down_pre.T @ x).to(lora_down.dtype) if needs_down else None + grad_lora_up = x_down.T @ grad_output if needs_up else None + return grad_x, None, None, None, None, grad_lora_down, grad_lora_up, None + + +class LinearGGUFA8( + GGUFLinear, + LoRAFusableLinearMixin, +): def __init__(self, dtype: torch.dtype, *args, **kwargs): super().__init__(*args, **kwargs) @@ -87,14 +153,38 @@ def __init__(self, dtype: torch.dtype, *args, **kwargs): def forward(self, x_orig: torch.Tensor) -> torch.Tensor: assert not self.weight.requires_grad x = x_orig.reshape(-1, x_orig.shape[-1]) - w = dequantize_gguf_tensor(self.weight.detach()) + w = dequantize_gguf_tensor(self.weight.detach()) + #the 8-bit path only pays for itself above a few tokens, and an unquantized GGUF type has + #nothing to requantize - dequantize_gguf_tensor already hands back the dense weight if x.shape[0] > 16 and self.weight.quant_type not in UNQUANTIZED_TYPES: - if self._dtype == torch.int8: - y = LinearGGUFIntA8RequantFunction.apply(x, w, self.bias, self.compute_dtype) - else: - y = LinearGGUFFpA8RequantFunction.apply(x, w, self.bias, self.compute_dtype) + #axiswise performs better than tensorwise in tests, even though + #it requires another requant during backward - but requant is cheap + y = LinearGGUFA8RequantFunction.apply(x, w, self._dtype, self.bias, self.compute_dtype) else: y = torch.nn.functional.linear(x, w, self.bias) return y.reshape(x_orig.shape[:-1] + (y.shape[-1], )) + + def forward_with_lora(self, x_orig: torch.Tensor, lora_down: torch.Tensor, lora_up: torch.Tensor, dropout: torch.nn.Dropout, alpha: torch.Tensor) -> torch.Tensor: + lora_rank = lora_down.shape[0] + + x = x_orig.reshape(-1, x_orig.shape[-1]) + if x.shape[0] <= 16 or self.weight.quant_type in UNQUANTIZED_TYPES: + ld = torch.nn.functional.linear(dropout(torch.nn.functional.linear(x_orig, lora_down)), lora_up) + return LinearGGUFA8.forward(self, x_orig) + ld * (alpha / lora_rank) + + #the cast matches lora_up (kept in its own dtype, e.g. f32) to the down-projection's autocast + #dtype, which the epilogue dot needs; dropout(ones) yields the scaled 0/(1/(1-p)) mask, drawn + #here rather than inside LinearGGUFA8RequantLoRAFunction so the RNG stays inductor-native + lora_up_scaled = (lora_up * (alpha / lora_rank)).T.to(self.compute_dtype) + dropout_mask = dropout(torch.ones(x.shape[0], lora_rank, device=x.device, dtype=self.compute_dtype)) if (dropout.training and dropout.p > 0) else None + y = self._fused_lora_forward(x, lora_down, lora_up_scaled, dropout_mask) + + return y.reshape(x_orig.shape[:-1] + (y.shape[-1], )) + + #see LinearW8A8._fused_lora_forward for what the fused node buys and what the operands are + def _fused_lora_forward(self, x: Tensor, lora_down: Tensor, lora_up: Tensor, dropout_mask: Tensor | None) -> Tensor: + assert not self.weight.requires_grad + w = dequantize_gguf_tensor(self.weight.detach()) + return LinearGGUFA8RequantLoRAFunction.apply(x, w, self._dtype, self.bias, self.compute_dtype, lora_down, lora_up, dropout_mask) diff --git a/modules/module/quantized/LinearSVD.py b/modules/module/quantized/LinearSVD.py index 16e2f2650..f352131c1 100644 --- a/modules/module/quantized/LinearSVD.py +++ b/modules/module/quantized/LinearSVD.py @@ -1,6 +1,6 @@ -from abc import abstractmethod from contextlib import suppress +from modules.module.quantized.mixin.LoRAFusableLinearMixin import LoRAFusableLinearMixin from modules.module.quantized.mixin.QuantizedLinearMixin import QuantizedLinearMixin from modules.module.quantized.mixin.QuantizedModuleMixin import QuantizedModuleMixin @@ -10,14 +10,11 @@ class BaseLinearSVD( QuantizedModuleMixin, QuantizedLinearMixin, + LoRAFusableLinearMixin, ): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - @abstractmethod - def forward_with_lora(self, x: torch.Tensor, lora_down: torch.nn.Linear, lora_up: torch.nn.Linear, dropout: torch.nn.Dropout, alpha: float) -> torch.Tensor: - pass - def _get_tensor_hash(t: torch.Tensor) -> str: t = t.flatten().to(torch.float32) vals = torch.stack([ @@ -105,20 +102,44 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: x_up = torch.nn.functional.linear(x_down, self.svd_up) return x_up + super().forward(x) - def forward_with_lora(self, x: torch.Tensor, lora_down: torch.nn.Linear, lora_up: torch.nn.Linear, dropout: torch.nn.Dropout, alpha: float) -> torch.Tensor: + def forward_with_lora(self, x: torch.Tensor, lora_down: torch.Tensor, lora_up: torch.Tensor, dropout: torch.nn.Dropout, alpha: torch.Tensor) -> torch.Tensor: assert self.__svd_is_quantized assert not self.svd_down.requires_grad and not self.svd_up.requires_grad - assert lora_down.bias is None and lora_up.bias is None - lora_rank = lora_down.weight.shape[0] - down_merged = torch.cat([lora_down.weight, self.svd_down], dim=0) + lora_rank = lora_down.shape[0] + flat = x.reshape(-1, x.shape[-1]) + from modules.module.quantized.LinearW8A8 import LinearW8A8 + #the fused kernel needs a LinearW8A8 base and enough tokens to beat the dense merge below + if flat.shape[0] > 16 and isinstance(self, LinearW8A8): + y = self.__fused_forward_with_lora(flat, lora_down, lora_up, dropout, alpha) + return y.reshape(x.shape[:-1] + (y.shape[-1], )) + + #autocast casts the merged matrix to the compute dtype for the matmul anyway, so cast the + #halves before the cat instead of building an f32 buffer to narrow again + down_merged = torch.cat([lora_down.to(self.compute_dtype), self.svd_down.to(self.compute_dtype)], dim=0) x_down = torch.nn.functional.linear(x, down_merged) if dropout.p > 0.0 and self.training: x_down[..., :lora_rank] = dropout(x_down[..., :lora_rank]) - lora_up_scaled = lora_up.weight * (alpha / lora_rank) - up_merged = torch.cat([lora_up_scaled, self.svd_up], dim=1) + lora_up_scaled = (lora_up * (alpha / lora_rank)).to(self.compute_dtype) + up_merged = torch.cat([lora_up_scaled, self.svd_up.to(self.compute_dtype)], dim=1) x_up = torch.nn.functional.linear(x_down, up_merged) return x_up + super().forward(x) + #fuse both low-rank branches into the residual mm's epilogue: concat the LoRA and SVD + #down/up into one rank r_lora + r_svd operand, alpha folded into the LoRA up-block only + #and a dropout mask that's 1 on the svd columns, so only LoRA drops. x is 2D (flattened) + def __fused_forward_with_lora(self, x: torch.Tensor, lora_down: torch.Tensor, lora_up: torch.Tensor, dropout: torch.nn.Dropout, alpha: torch.Tensor) -> torch.Tensor: + lora_rank = lora_down.shape[0] + + down_merged = torch.cat([lora_down.to(self.compute_dtype), self.svd_down.to(self.compute_dtype)], dim=0) + up_fused = torch.cat([(lora_up * (alpha / lora_rank)).to(self.compute_dtype), self.svd_up.to(self.compute_dtype)], dim=1).T + dropout_mask = None + if dropout.training and dropout.p > 0: + svd_rank = self.svd_down.shape[0] + lora_mask = dropout(torch.ones(x.shape[0], lora_rank, device=x.device, dtype=self.compute_dtype)) + dropout_mask = torch.cat([lora_mask, torch.ones(x.shape[0], svd_rank, device=x.device, dtype=self.compute_dtype)], dim=1) + + return self._fused_lora_forward(x, down_merged, up_fused, dropout_mask) + return LinearSVD diff --git a/modules/module/quantized/LinearW8A8.py b/modules/module/quantized/LinearW8A8.py index 8bbb90e5a..a49b7c9f7 100644 --- a/modules/module/quantized/LinearW8A8.py +++ b/modules/module/quantized/LinearW8A8.py @@ -1,10 +1,12 @@ - from modules.module.quantized.mixin.CompressedWeightMixin import CompressedWeightMixin +from modules.module.quantized.mixin.LoRAFusableLinearMixin import LoRAFusableLinearMixin from modules.module.quantized.mixin.QuantizedLinearMixin import QuantizedLinearMixin from modules.module.quantized.mixin.QuantizedModuleMixin import QuantizedModuleMixin from modules.util.mm_8bit import mm_8bit as mm_8bit +from modules.util.mm_8bit import scaled_lora_mm_8bit, scaled_mm_8bit from modules.util.quantization_util import ( dequantize, + quantize_axiswise, quantize_fp8_axiswise, quantize_fp8_tensorwise, quantize_int8_axiswise, @@ -16,40 +18,64 @@ @torch.no_grad() -def int8_forward_tokenwise(x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: - x_8, x_scale = quantize_int8_axiswise(x, dim=-1) - res = torch._int_mm(x_8, weight.T) - res_scaled = res.float().mul_(weight_scale * x_scale).to(compute_dtype) +def forward_tokenwise_postscaled_torch(x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: + if weight.dtype == torch.int8: + x_8, x_scale = quantize_int8_axiswise(x, dim=-1) + res = torch._int_mm(x_8, weight.T) + res_scaled = res.float().mul_(weight_scale * x_scale).to(compute_dtype) + else: + x_8, x_scale = quantize_fp8_axiswise(x, dim=-1) + one = torch.tensor(1.0, device=x.device) + res = torch._scaled_mm(x_8, weight.T, scale_a=one, scale_b=weight_scale.float(), out_dtype=torch.float) + res_scaled = res.mul_(x_scale).to(compute_dtype) #much faster than scaled by _scaled_mm if bias is not None: res_scaled.add_(bias) return res_scaled @torch.no_grad() -def fp8_forward_tokenwise(x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: - x_8, x_scale = quantize_fp8_axiswise(x, dim=-1) - one = torch.tensor(1.0, device=x.device) - res = torch._scaled_mm(x_8, weight.T, scale_a=one, scale_b=weight_scale.float(), out_dtype=torch.float) - res_scaled = res.mul_(x_scale).to(compute_dtype) #much faster than scaled by _scaled_mm +def forward_tokenwise_epiloguescaled_triton(x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: + x_8, x_scale = quantize_axiswise(x, dim=-1, dtype=weight.dtype) + res_scaled = scaled_mm_8bit(x_8, weight.T, weight_scale * x_scale, compute_dtype) if bias is not None: res_scaled.add_(bias) return res_scaled @torch.no_grad() -def int8_backward_axiswise(output: Tensor, weight: Tensor, weight_scale: Tensor) -> Tensor: - output_8, output_scale = quantize_int8_axiswise(output, dim=-1) +def backward_tokenwise_postscaled_triton(output: Tensor, weight: Tensor, weight_scale: Tensor) -> Tensor: + output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=weight.dtype) #almost always, grad outputs are already contiguous and this is a no-op. But there are some grad outputs from SDXL that are non-contiguous: mm_res = mm_8bit(output_8.contiguous(), weight) return mm_res.float().mul_(weight_scale * output_scale).to(output.dtype) @torch.no_grad() -def fp8_backward_axiswise(output: Tensor, weight: Tensor, weight_scale: Tensor) -> Tensor: - output_8, output_scale = quantize_fp8_axiswise(output, dim=-1) - mm_res = mm_8bit(output_8.contiguous(), weight) - return mm_res.float().mul_(weight_scale * output_scale).to(output.dtype) +def backward_tokenwise_epiloguescaled_triton(output: Tensor, weight: Tensor, weight_scale: Tensor) -> Tensor: + output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=weight.dtype) + return scaled_mm_8bit(output_8.contiguous(), weight, weight_scale * output_scale, output.dtype) + + +forward_tokenwise = forward_tokenwise_epiloguescaled_triton +backward_tokenwise = backward_tokenwise_epiloguescaled_triton + + +@torch.no_grad() +def forward_tokenwise_lora_epiloguescaled_triton(x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype, x_down: Tensor, lora_up: Tensor) -> Tensor: + x_8, x_scale = quantize_axiswise(x, dim=-1, dtype=weight.dtype) + res_scaled = scaled_lora_mm_8bit(x_8, weight.T, weight_scale * x_scale, x_down, lora_up, compute_dtype) + if bias is not None: + res_scaled.add_(bias) + return res_scaled + +@torch.no_grad() +def backward_tokenwise_lora_epiloguescaled_triton(output: Tensor, weight: Tensor, weight_scale: Tensor, grad_x_down_pre: Tensor, lora_down: Tensor) -> Tensor: + output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=weight.dtype) + return scaled_lora_mm_8bit(output_8.contiguous(), weight, weight_scale * output_scale, grad_x_down_pre, lora_down.to(grad_x_down_pre.dtype), output.dtype) + + +forward_tokenwise_lora = forward_tokenwise_lora_epiloguescaled_triton +backward_tokenwise_lora = backward_tokenwise_lora_epiloguescaled_triton class LinearW8A8Function(torch.autograd.Function): - #int8 and fp8 differ only in which math runs, and the quantized weight already carries the dtype @staticmethod def forward(ctx, x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: # `weight` is the decompressed weight, so saving it keeps a full-size copy alive until backward. @@ -59,10 +85,7 @@ def forward(ctx, x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | # TODO once offloading uses non-reentrant checkpointing, consider not saving the decompressed # weight and decoding in backward() instead. ctx.save_for_backward(weight, weight_scale) - if weight.dtype == torch.int8: - return int8_forward_tokenwise(x, weight, weight_scale, bias, compute_dtype) - else: - return fp8_forward_tokenwise(x, weight, weight_scale, bias, compute_dtype) + return forward_tokenwise(x, weight, weight_scale, bias, compute_dtype) @staticmethod def backward(ctx, output: Tensor): @@ -70,10 +93,40 @@ def backward(ctx, output: Tensor): raise NotImplementedError("Int/Float A8W8 cannot be used for full finetuning") weight, weight_scale = ctx.saved_tensors - if weight.dtype == torch.int8: - return int8_backward_axiswise(output, weight, weight_scale), None, None, None, None - else: - return fp8_backward_axiswise(output, weight, weight_scale), None, None, None, None + return backward_tokenwise(output, weight, weight_scale), None, None, None, None + + +class LinearW8A8LoRAFunction(torch.autograd.Function): + #LinearW8A8Function plus a low-rank update: it owns the down-projection and the dropout, so the LoRA + #dgrad folds into the backward epilogue and the (M, out) product is never materialized. The weight is + #the already decompressed one, as in LinearW8A8Function + @staticmethod + def forward(ctx, x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype, lora_down: Tensor, lora_up: Tensor, dropout_mask: Tensor | None) -> Tensor: + x_down = torch.nn.functional.linear(x, lora_down) + if dropout_mask is not None: + x_down = x_down * dropout_mask + #backward runs outside autocast, so x is saved in x_down's dtype: grad_lora_down is a + #compute-dtype product there, as it is in the unfused path where autocast casts x + ctx.save_for_backward(weight, weight_scale, x.to(x_down.dtype), lora_down, lora_up, dropout_mask, x_down) + return forward_tokenwise_lora(x, weight, weight_scale, bias, compute_dtype, x_down, lora_up) + + @staticmethod + def backward(ctx, grad_output: Tensor): + if ctx.needs_input_grad[1:5] != (False, False, False, False): + raise NotImplementedError("Int/Float A8W8 cannot be used for full finetuning") + + weight, weight_scale, x, lora_down, lora_up, dropout_mask, x_down = ctx.saved_tensors + needs_x, needs_down, needs_up = ctx.needs_input_grad[0], ctx.needs_input_grad[5], ctx.needs_input_grad[6] + #grad_x_down_pre (post-dropout-backward) feeds both the grad_x epilogue fold and grad_lora_down + grad_x_down_pre = None + if needs_x or needs_down: + grad_x_down_pre = grad_output @ lora_up.T + if dropout_mask is not None: + grad_x_down_pre = grad_x_down_pre * dropout_mask + grad_x = backward_tokenwise_lora(grad_output, weight, weight_scale, grad_x_down_pre, lora_down) if needs_x else None + grad_lora_down = (grad_x_down_pre.T @ x).to(lora_down.dtype) if needs_down else None + grad_lora_up = x_down.T @ grad_output if needs_up else None + return grad_x, None, None, None, None, grad_lora_down, grad_lora_up, None class LinearW8A8( @@ -81,6 +134,7 @@ class LinearW8A8( QuantizedModuleMixin, QuantizedLinearMixin, CompressedWeightMixin, + LoRAFusableLinearMixin, ): def __init__(self, dtype: torch.dtype, *args, **kwargs): super().__init__(*args, **kwargs) @@ -145,6 +199,30 @@ def forward(self, x_orig: torch.Tensor) -> torch.Tensor: return y.reshape(x_orig.shape[:-1] + (y.shape[-1], )) + def forward_with_lora(self, x_orig: torch.Tensor, lora_down: torch.Tensor, lora_up: torch.Tensor, dropout: torch.nn.Dropout, alpha: torch.Tensor) -> torch.Tensor: + assert self.__is_quantized + lora_rank = lora_down.shape[0] + + x = x_orig.reshape(-1, x_orig.shape[-1]) + if x.shape[0] <= 16: + ld = torch.nn.functional.linear(dropout(torch.nn.functional.linear(x_orig, lora_down)), lora_up) + return LinearW8A8.forward(self, x_orig) + ld * (alpha / lora_rank) + + #the cast matches lora_up (kept in its own dtype, e.g. f32) to the down-projection's autocast + #dtype, which the epilogue dot needs; dropout(ones) yields the scaled 0/(1/(1-p)) mask, drawn + #here rather than inside LinearW8A8LoRAFunction so the RNG stays inductor-native + lora_up_scaled = (lora_up * (alpha / lora_rank)).T.to(self.compute_dtype) + dropout_mask = dropout(torch.ones(x.shape[0], lora_rank, device=x.device, dtype=self.compute_dtype)) if (dropout.training and dropout.p > 0) else None + y = self._fused_lora_forward(x, lora_down, lora_up_scaled, dropout_mask) + + return y.reshape(x_orig.shape[:-1] + (y.shape[-1], )) + + #x is 2D, lora_down is (r, in_features), lora_up is (r, out_features) with alpha already folded in + def _fused_lora_forward(self, x: Tensor, lora_down: Tensor, lora_up: Tensor, dropout_mask: Tensor | None) -> Tensor: + assert not self.weight.requires_grad + weight = self._decompress(self.weight.detach()) if self._compressed else self.weight + return LinearW8A8LoRAFunction.apply(x, weight, self.scale, self.bias, self.compute_dtype, lora_down, lora_up, dropout_mask) + def run_benchmark(fn, desc, steps=10000, warmup=500, compile=False): if compile: fn = torch.compile(fn, fullgraph=True) @@ -158,7 +236,7 @@ def run_benchmark(fn, desc, steps=10000, warmup=500, compile=False): @torch.no_grad() -def benchmark_int8(m, k, n, device = 'cuda'): +def benchmark_int8(m, k, n, device = 'cuda', steps = 10000): x = torch.randn(m,k, device=device, dtype=torch.bfloat16) x_8 = torch.ones (m,k, device=device, dtype=torch.int8) y = torch.randn(m,n, device=device, dtype=torch.bfloat16) @@ -167,19 +245,21 @@ def benchmark_int8(m, k, n, device = 'cuda'): w_scale = torch.ones(1, device=device) - run_benchmark(lambda: torch._int_mm(x_8, w_8.T), "torch mm int") - run_benchmark(lambda: mm_8bit(x_8, w_8.T), "triton mm int") + run_benchmark(lambda: torch._int_mm(x_8, w_8.T), "torch mm int", steps=steps) + run_benchmark(lambda: mm_8bit(x_8, w_8.T), "triton mm int", steps=steps) def torch_backward(a, b): torch._int_mm(a, b.T.contiguous().T) - run_benchmark(lambda: torch_backward(y_8, w_8), "torch mm backward int8") - run_benchmark(lambda: mm_8bit(y_8, w_8), "triton mm backward int8") + run_benchmark(lambda: torch_backward(y_8, w_8), "torch mm backward int8", steps=steps) + run_benchmark(lambda: mm_8bit(y_8, w_8), "triton mm backward int8", steps=steps) - run_benchmark(lambda: int8_forward_tokenwise(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "torch forward int", compile=True) - run_benchmark(lambda: int8_backward_axiswise(y, w_8, w_scale), "triton backward int", compile=True) + run_benchmark(lambda: forward_tokenwise_postscaled_torch(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "torch forward int", steps=steps, compile=True) + run_benchmark(lambda: backward_tokenwise_postscaled_triton(y, w_8, w_scale), "triton backward int", steps=steps, compile=True) + run_benchmark(lambda: forward_tokenwise_epiloguescaled_triton(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "triton scaled forward int", steps=steps, compile=True) + run_benchmark(lambda: backward_tokenwise_epiloguescaled_triton(y, w_8, w_scale), "triton scaled backward int", steps=steps, compile=True) @torch.no_grad() -def benchmark_fp8(m, k, n, device = 'cuda'): +def benchmark_fp8(m, k, n, device = 'cuda', steps = 10000): x = torch.randn(m,k, device=device, dtype=torch.bfloat16) x_8 = torch.ones (m,k, device=device, dtype=torch.float8_e4m3fn) y = torch.randn(m,n, device=device, dtype=torch.bfloat16) @@ -188,16 +268,93 @@ def benchmark_fp8(m, k, n, device = 'cuda'): w_scale = torch.ones(1, device=device, dtype=torch.bfloat16) one_scale = torch.ones(1, device=device) - run_benchmark(lambda: torch._scaled_mm(x_8, w_8.T, out_dtype=torch.bfloat16, scale_a=one_scale.float(), scale_b=w_scale.float()), "torch mm fp8") - run_benchmark(lambda: mm_8bit(x_8, w_8.T), "triton mm fp8") + run_benchmark(lambda: torch._scaled_mm(x_8, w_8.T, out_dtype=torch.bfloat16, scale_a=one_scale.float(), scale_b=w_scale.float()), "torch mm fp8", steps=steps) + run_benchmark(lambda: mm_8bit(x_8, w_8.T), "triton mm fp8", steps=steps) def torch_backward(a, b): torch._scaled_mm(a, b.T.contiguous().T, out_dtype=torch.bfloat16, scale_a=one_scale.float(), scale_b=w_scale.float()) - run_benchmark(lambda: torch_backward(y_8, w_8), "torch mm backward fp8") - run_benchmark(lambda: mm_8bit(y_8, w_8), "triton mm backward fp8") - run_benchmark(lambda: fp8_forward_tokenwise(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "torch forward fp8", compile=True) - run_benchmark(lambda: fp8_backward_axiswise(y, w_8, w_scale), "triton backward fp8", compile=True) + run_benchmark(lambda: torch_backward(y_8, w_8), "torch mm backward fp8", steps=steps) + run_benchmark(lambda: mm_8bit(y_8, w_8), "triton mm backward fp8", steps=steps) + + run_benchmark(lambda: forward_tokenwise_postscaled_torch(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "torch forward fp8", steps=steps, compile=True) + run_benchmark(lambda: forward_tokenwise_epiloguescaled_triton(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "triton scaled forward fp8", steps=steps, compile=True) + run_benchmark(lambda: backward_tokenwise_epiloguescaled_triton(y, w_8, w_scale), "triton scaled backward fp8", steps=steps, compile=True) + + +@torch.no_grad() +def benchmark_lora(m, k, n, r, dtype, device='cuda', steps=1000): + is_int8 = (dtype == torch.int8) + x = torch.randn(m, k, device=device, dtype=torch.bfloat16) + if is_int8: + w_8 = torch.randint(-127, 127, (n, k), device=device, dtype=torch.int8) + else: + w_8 = torch.randn(n, k, device=device).to(torch.float8_e4m3fn) + w_scale = torch.full((1,), 0.01, device=device) + down_w = torch.randn(r, k, device=device, dtype=torch.bfloat16) * 0.02 + up_w = torch.randn(n, r, device=device, dtype=torch.bfloat16) * 0.02 + + baseline = forward_tokenwise + fused = forward_tokenwise_lora_epiloguescaled_triton + + def run_unfused(): + y = baseline(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16) + return y + (x @ down_w.T) @ up_w.T + + def run_fused(): + x_down = x @ down_w.T + return fused(x, w_8, w_scale, None, torch.bfloat16, x_down, up_w.T) + + name = "int8" if is_int8 else "fp8" + diff = (run_unfused().float() - run_fused().float()).abs().max().item() + ref = run_unfused().float().abs().mean().item() + print(f"lora {name} m={m} k={k} n={n} r={r}: max abs diff fused vs unfused = {diff:.4f} (mean magnitude {ref:.3f})") + run_benchmark(lambda: baseline(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), f"no-lora {name} baseline", steps=steps, compile=True) + run_benchmark(run_unfused, f"unfused lora {name}", steps=steps, compile=True) + run_benchmark(run_fused, f"fused lora {name}", steps=steps, compile=True) + + +@torch.no_grad() +def benchmark_lora_backward(m, k, n, r, dtype, device='cuda', steps=1000): + #isolates the grad_x path that the backward fusion changes: unfused = main bwd mm + + #standalone (M,K) lora dgrad + merge add; fused = the same lora dgrad folded into the + #bwd mm epilogue. grad_x_down (= grad_y @ up.T) is shared and timed in both; the + #grad_lora_up / grad_lora_down gemms are identical in both variants and excluded. + is_int8 = (dtype == torch.int8) + grad_y = torch.randn(m, n, device=device, dtype=torch.bfloat16) + if is_int8: + w_8 = torch.randint(-127, 127, (n, k), device=device, dtype=torch.int8) + else: + w_8 = torch.randn(n, k, device=device).to(torch.float8_e4m3fn) + w_scale = torch.full((1,), 0.01, device=device) + down_w = torch.randn(r, k, device=device, dtype=torch.bfloat16) * 0.02 + up_w = torch.randn(n, r, device=device, dtype=torch.bfloat16) * 0.02 + + main_bwd = backward_tokenwise + fused_bwd = backward_tokenwise_lora + + def run_unfused(): + grad_x = main_bwd(grad_y, w_8, w_scale) + grad_x_down = grad_y @ up_w + return grad_x + grad_x_down @ down_w + + def run_fused(): + grad_x_down = grad_y @ up_w + return fused_bwd(grad_y, w_8, w_scale, grad_x_down, down_w) + + name = "int8" if is_int8 else "fp8" + diff = (run_unfused().float() - run_fused().float()).abs().max().item() + ref = run_unfused().float().abs().mean().item() + print(f"lora-bwd {name} m={m} k={k} n={n} r={r}: max abs diff fused vs unfused = {diff:.4f} (mean magnitude {ref:.3f})") + run_benchmark(run_unfused, f"unfused lora-bwd {name}", steps=steps, compile=True) + run_benchmark(run_fused, f"fused lora-bwd {name}", steps=steps, compile=True) if __name__ == "__main__": + #ragged shape: M and the backward's contraction over 3088 are both non-%128, so the + #masked-loop fallback and, with r=12, the padded rank tile are covered benchmark_int8(2 * 1024 + 50, 3072, 3072 + 16) benchmark_fp8(2 * 1024 + 50, 3072, 3072 + 16) + benchmark_lora(2 * 1024 + 50, 3072, 3072 + 16, r=12, dtype=torch.int8) + benchmark_lora_backward(2 * 1024 + 50, 3072, 3072 + 16, r=12, dtype=torch.int8) + #a real FLUX.2-klein-9B shape at 512px batch 2: M = 2*(1024 image + 512 text) tokens + #through a double block's 4096 -> 4096 attention projection + benchmark_lora(2 * (1024 + 512), 4096, 4096, r=16, dtype=torch.int8, steps=1000) diff --git a/modules/module/quantized/mixin/LoRAFusableLinearMixin.py b/modules/module/quantized/mixin/LoRAFusableLinearMixin.py new file mode 100644 index 000000000..784b9de0d --- /dev/null +++ b/modules/module/quantized/mixin/LoRAFusableLinearMixin.py @@ -0,0 +1,9 @@ +from abc import ABCMeta, abstractmethod + +import torch + + +class LoRAFusableLinearMixin(metaclass=ABCMeta): + @abstractmethod + def forward_with_lora(self, x: torch.Tensor, lora_down: torch.Tensor, lora_up: torch.Tensor, dropout: torch.nn.Dropout, alpha: torch.Tensor) -> torch.Tensor: + pass diff --git a/modules/util/mm_8bit.py b/modules/util/mm_8bit.py index a4e74f4bf..e76a2155e 100644 --- a/modules/util/mm_8bit.py +++ b/modules/util/mm_8bit.py @@ -1,8 +1,15 @@ +import torch + try: - from modules.util.triton_mm_8bit import mm_8bit + from modules.util.triton_mm_8bit import ( + mm_8bit, + rowcol_scaled_lora_mm_8bit, + rowcol_scaled_mm_8bit, + scaled_lora_mm_8bit, + scaled_mm_8bit, + ) except ImportError as e: print(str(e) + ", continuing without triton") - import torch def mm_8bit(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: assert a.shape[1] == b.shape[0], "Incompatible dimensions" assert a.is_contiguous(), "Matrix A must be contiguous" @@ -12,4 +19,17 @@ def mm_8bit(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: return torch._int_mm(a, b) else: one = torch.ones(1, device=a.device) - return torch._scaled_mm(a, b.T.contiguous().T, scale_a=one, scale_b=one) + #out_dtype defaults to a's dtype, which would round the accumulator back to fp8 + return torch._scaled_mm(a, b.T.contiguous().T, scale_a=one, scale_b=one, out_dtype=torch.float32) + + def scaled_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return mm_8bit(a, b).float().mul_(scale.reshape(-1, 1)).to(out_dtype) + + def scaled_lora_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return mm_8bit(a, b).float().mul_(scale.reshape(-1, 1)).add_(xd @ up).to(out_dtype) + + def rowcol_scaled_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return mm_8bit(a, b).float().mul_(scale.reshape(-1, 1)).mul_(scale_n.reshape(1, -1)).to(out_dtype) + + def rowcol_scaled_lora_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return mm_8bit(a, b).float().mul_(scale.reshape(-1, 1)).mul_(scale_n.reshape(1, -1)).add_(xd @ up).to(out_dtype) diff --git a/modules/util/quantization_util.py b/modules/util/quantization_util.py index d7229687d..e156fb855 100644 --- a/modules/util/quantization_util.py +++ b/modules/util/quantization_util.py @@ -74,6 +74,14 @@ def quantize_fp8_axiswise(x: Tensor, dim: int) -> tuple[Tensor, Tensor]: q = quantize_fp8(x, scale) return q, scale +def quantize_axiswise(x: Tensor, dim: int, dtype: torch.dtype) -> tuple[Tensor, Tensor]: + if dtype == torch.int8: + return quantize_int8_axiswise(x, dim) + elif dtype == torch.float8_e4m3fn: + return quantize_fp8_axiswise(x, dim) + else: + raise NotImplementedError(f"{dtype} is not an 8-bit quantization dtype") + def dequantize(q: Tensor, scale: float | Tensor) -> Tensor: return q.float() * scale diff --git a/modules/util/triton_mm_8bit.py b/modules/util/triton_mm_8bit.py index 522959d89..62d26dc50 100644 --- a/modules/util/triton_mm_8bit.py +++ b/modules/util/triton_mm_8bit.py @@ -1,74 +1,155 @@ -#This is a 8bit matmul kernel adapted from the Triton tutorial here: +#8bit matmul kernels adapted from the Triton tutorial here: #https://triton-lang.org/main/getting-started/tutorials/03-matrix-multiplication.html -#It is not optimized and about 10% slower than torch._int_mm and torch._scaled_mm -#However, the torch functions don't work on row-major rhs matrices: -#_scaled_mm fails, _int_mm automatically converts to column-major -# -#Converting to column-major is slow, which is significant because the weights matrix -#of a Linear layer is always column-major during the backward pass. -# -#In these cases, this Triton kernel is *much* faster because it can access the -#row-major weight matrix directly, using strided memory access +#All three mm entry points share the _mm_accumulate compute core (grouped launch order +#for L2 reuse, compile-time layout/divisibility specialization, EVEN_K loop selection) +#and differ only in the epilogue: _mm_kernel stores the raw int32/fp32 product, +#_scaled_mm_kernel folds a per-row dequant scale in and casts to the output dtype, and +#_scaled_lora_mm_kernel additionally fuses a low-rank update into the same tile. import torch import triton import triton.language as tl +from tqdm import tqdm -@triton.autotune( - configs=[ - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 128}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 128}, num_stages=3,num_warps=8), - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64}, num_stages=3,num_warps=8), - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 64}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 64}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 128}, num_stages=3,num_warps=8), - triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 128}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=5,num_warps=2), - triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 128}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32}, num_stages=4,num_warps=4), - triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=5,num_warps=2), - - - ], - key=[ - 'QUANTIZED_M', #only tune roughly on M, because M is the transformer sequence length - can vary on data - 'N', - 'K', - 'stride_bk' #use stride of b as key, to autotune again for a strided rhs matrix (backward pass) - ], -) +def announce_autotuning(kernel, name=None): + prefix = f"[triton] Autotuning {name} " if name else "[triton] Autotuning " + orig_check_disk_cache = kernel.check_disk_cache + variants = 0 + def check_disk_cache(tuning_key, configs, bench_fn): + def announced_bench(): + nonlocal variants + variants += 1 + tqdm.write(f"{prefix}variant #{variants}...") + bench_fn() + return orig_check_disk_cache(tuning_key, configs, announced_bench) + kernel.check_disk_cache = check_disk_cache + +#tiled 8-bit transpose, used to rewrite the backward's B matrix to k-major before the mm. +#The int8 tensor-core op needs its B operand k-major and Ada has no 8-bit ldmatrix.trans, so an +#n-major B - a Linear weight in the backward pass - makes the mm emulate the transpose with byte +#shuffles in its inner loop, ~30% slower on every shape. This copy runs at memory bandwidth +#(~570GB/s) once per weight instead, and the mm then takes the fast k-major path. +_TRANSPOSE_AUTOTUNE_CONFIGS = [ + triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 128}, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 64}, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128}, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128}, num_warps=8), +] + +@triton.autotune(configs=_TRANSPOSE_AUTOTUNE_CONFIGS, key=['M', 'N'], cache_results=True) @triton.jit -def _mm_kernel( - a_ptr, b_ptr, c_ptr, +def _transpose_kernel( + src_ptr, dst_ptr, + M, N, + stride_sm, stride_dn, + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, +): + pid_m = tl.program_id(axis=0) + pid_n = tl.program_id(axis=1) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + tile = tl.load(src_ptr + offs_m[:, None] * stride_sm + offs_n[None, :], + mask=(offs_m[:, None] < M) & (offs_n[None, :] < N)) + tl.store(dst_ptr + offs_n[:, None] * stride_dn + offs_m[None, :], + tl.trans(tile), + mask=(offs_n[:, None] < N) & (offs_m[None, :] < M)) + +def transpose_8bit(src: torch.Tensor) -> torch.Tensor: + #returns src^T as a new contiguous tensor (any 1-byte dtype) + assert src.stride(1) == 1, "src must be contiguous along axis 1" + M, N = src.shape + dst = torch.empty((N, M), device=src.device, dtype=src.dtype) + + def grid(META): + return (triton.cdiv(M, META['BLOCK_SIZE_M']), triton.cdiv(N, META['BLOCK_SIZE_N'])) + _transpose_kernel[grid]( + src, dst, + M, N, + src.stride(0), dst.stride(0), + ) + return dst + +announce_autotuning(_transpose_kernel, name="8-bit transpose") + +#minimum M for the transpose-to-k-major rewrite in the mm wrappers below: the mm saves ~30% +#(~205 -> ~275 TOPS) but the copy costs 2*K*N bytes of traffic regardless of M, so the rewrite +#only pays above a token count. Breakeven is shape-dependent - measured on Ada (4070 Ti SUPER) +#at ~1550 for the widest layers and below 512 for the attention projections - and the wrappers +#see only M, so this takes the widest layer's breakeven and every shape wins above it +_TRANSPOSE_MIN_M = 1536 + + +_AUTOTUNE_KEY = [ + 'QUANTIZED_M', #only tune roughly on M, because M is the transformer sequence length - can vary on data + 'N', + 'K', + 'stride_bk' #use stride of b as key, to autotune again for a strided rhs matrix (backward pass) +] + +#the LoRA epilogue keys on the rank tiling as well, and on stride_upr for the same reason as +#stride_bk above: up arrives row-major (r, N) when the caller's cast to the compute dtype is a +#real copy, and as a transposed view when it is a no-op, and the two layouts load differently +_LORA_AUTOTUNE_KEY = [*_AUTOTUNE_KEY, 'R_TILES', 'BLOCK_R', 'stride_upr'] + +#configs for the shared _mm_accumulate core: GROUP_SIZE_M is required by its grouped launch +#order. Shared memory per config is stages*(BLOCK_M+BLOCK_N)*BLOCK_K bytes and must stay +#under the ~99KB per-CTA limit (identical on sm86/sm89/sm120, i.e. consumer/workstation +#Ampere through Blackwell); oversized configs would be skipped by the autotuner. +#GROUP_SIZE_M=8 suits the large L2 of Ada/Blackwell; the 16-variants keep fewer B +#panels in flight per wave, for the small L2 (4-6MB) of Ampere - autotuning picks +_AUTOTUNE_CONFIGS = [ + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8}, num_stages=4, num_warps=8), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 16}, num_stages=4, num_warps=8), + triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8}, num_stages=4, num_warps=8), + triton.Config({'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 16}, num_stages=4, num_warps=8), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 8}, num_stages=3, num_warps=8), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 16}, num_stages=3, num_warps=8), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 8}, num_stages=3, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8}, num_stages=4, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8}, num_stages=4, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 8}, num_stages=4, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 64, 'GROUP_SIZE_M': 8}, num_stages=5, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 8}, num_stages=5, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 8}, num_stages=4, num_warps=4), + triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 128, 'GROUP_SIZE_M': 8}, num_stages=5, num_warps=2), +] + + +#shared compute core of the 8-bit mm kernels: grouped launch order, divisibility hints and +#the EVEN_K-specialized main loop. Returns the raw accumulator plus the output tile offsets; +#each entry kernel below adds its own epilogue and stores. +@triton.jit +def _mm_accumulate( + a_ptr, b_ptr, M, N, K, - stride_am, stride_ak, - stride_bk, stride_bn, - stride_cm, stride_cn, - BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, - QUANTIZED_M, - FLOAT: tl.constexpr, + stride_am, stride_ak, stride_bk, stride_bn, + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, + FLOAT: tl.constexpr, EVEN_K: tl.constexpr, ): - pid_n = tl.program_id(axis=0) - pid_m = tl.program_id(axis=1) + #grouped launch order: consecutive pids walk down GROUP_SIZE_M M-blocks before advancing + #to the next N-block, so the concurrent wave covers a rectangle of blocks and each B panel + #is read from DRAM once and reused by GROUP_SIZE_M CTAs out of L2. A naive row-major grid + #re-reads all of B per M-row instead, which makes the mm DRAM-bound + pid = tl.program_id(axis=0) + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m tl.assume(pid_m >= 0) tl.assume(pid_n >= 0) tl.assume(stride_am > 0) - tl.assume(stride_ak > 0) tl.assume(stride_bn > 0) tl.assume(stride_bk > 0) - tl.assume(stride_cm > 0) - tl.assume(stride_cn > 0) offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N @@ -78,24 +159,41 @@ def _mm_kernel( accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32 if FLOAT else tl.int32) - for k in range(tl.cdiv(K, BLOCK_SIZE_K)): - a_mask = (offs_am[:, None] < M) & (offs_k[None, :] < K - k*BLOCK_SIZE_K) - b_mask = (offs_bn[None, :] < N) & (offs_k[:, None] < K - k*BLOCK_SIZE_K) - a = tl.load(a_ptrs, mask=a_mask, other=0.0) - b = tl.load(b_ptrs, mask=b_mask, other=0.0) + if EVEN_K: + for _k in range(K // BLOCK_SIZE_K): + a = tl.load(a_ptrs) + b = tl.load(b_ptrs) - accumulator = tl.dot(a, b, accumulator, out_dtype=tl.float32 if FLOAT else tl.int32) + accumulator = tl.dot(a, b, accumulator, out_dtype=tl.float32 if FLOAT else tl.int32) - a_ptrs += BLOCK_SIZE_K * stride_ak - b_ptrs += BLOCK_SIZE_K * stride_bk + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk + else: + for k in range(tl.cdiv(K, BLOCK_SIZE_K)): + a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k*BLOCK_SIZE_K, other=0.0) + b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k*BLOCK_SIZE_K, other=0.0) + + accumulator = tl.dot(a, b, accumulator, out_dtype=tl.float32 if FLOAT else tl.int32) + + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) - c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + return accumulator, offs_cm, offs_cn + + +#shared epilogue tail for every mm kernel: c is allocated row-major in _prepare_mm, so only +#stride_cm is passed and the n stride is 1 +@triton.jit +def _store_c(c_ptr, value, offs_cm, offs_cn, M, N, stride_cm): + tl.assume(stride_cm > 0) + c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + offs_cn[None, :] c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N) - tl.store(c_ptrs, accumulator, mask=c_mask) + tl.store(c_ptrs, value, mask=c_mask) -def mm_8bit(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + +def _prepare_mm(a: torch.Tensor, b: torch.Tensor, out_dtype: torch.dtype): assert a.shape[1] == b.shape[0], "Incompatible dimensions" assert a.is_contiguous(), "Matrix A must be contiguous" assert a.dtype == b.dtype, "Incompatible dtypes" @@ -105,17 +203,294 @@ def mm_8bit(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: M, K = a.shape K, N = b.shape - c = torch.empty((M, N), device=a.device, dtype=torch.float32 if FLOAT else torch.int32) + #the kernel handles exactly two B layouts: k-major (forward, weight.T) and n-major (backward, weight) + B_K_MAJOR = (b.stride(0) == 1) + assert B_K_MAJOR or b.stride(1) == 1, "Matrix B must be contiguous along one axis" + #n-major B runs the mm ~30% slower; for large M a transpose copy pays for itself (see transpose_8bit) + if not B_K_MAJOR and M >= _TRANSPOSE_MIN_M: + b = transpose_8bit(b).t() + B_K_MAJOR = True + c = torch.empty((M, N), device=a.device, dtype=out_dtype) + return b, c, M, N, K, FLOAT + + +def _prepare_scaled_mm(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, out_dtype: torch.dtype): + #_prepare_mm plus the per-row scale reshape (the scaled kernels fold it into the epilogue) + b, c, M, N, K, FLOAT = _prepare_mm(a, b, out_dtype) + scale = scale.reshape(-1).to(torch.float32).contiguous() + assert scale.shape[0] == M, "scale must have one entry per row of a" + return b, c, scale, M, N, K, FLOAT + + +def _prepare_rowcol_scaled_mm(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, out_dtype: torch.dtype): + b, c, scale, M, N, K, FLOAT = _prepare_scaled_mm(a, b, scale, out_dtype) + scale_n = scale_n.reshape(-1).to(torch.float32).contiguous() + assert scale_n.shape[0] == N, "scale_n must have one entry per column of b" + return b, c, scale, scale_n, M, N, K, FLOAT + + +@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_AUTOTUNE_KEY, cache_results=True) +@triton.jit +def _mm_kernel( + a_ptr, b_ptr, c_ptr, + M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, +): + accumulator, offs_cm, offs_cn = _mm_accumulate( + a_ptr, b_ptr, M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, + BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, + FLOAT, EVEN_K, + ) + + _store_c(c_ptr, accumulator, offs_cm, offs_cn, M, N, stride_cm) + +announce_autotuning(_mm_kernel, name="8-bit matmul") + +#Opaque custom ops work around pytorch#164124: torch.compile otherwise absorbs the traced +#@triton.autotune kernels and freezes the config benchmarked for the first shape. Kept opaque, +#these bodies run eagerly, so Triton's autotuner selects per key and the JIT sees real sizes +@torch.library.custom_op("ot_quant::mm_8bit", mutates_args=()) +def mm_8bit(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + out_dtype = torch.float32 if a.dtype == torch.float8_e4m3fn else torch.int32 + b, c, M, N, K, FLOAT = _prepare_mm(a, b, out_dtype) + + #1D grid: the kernel derives pid_m/pid_n itself in grouped order for L2 reuse def grid(META): - return (triton.cdiv(N, META['BLOCK_SIZE_N']) , triton.cdiv(M, META['BLOCK_SIZE_M']), ) + return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) _mm_kernel[grid]( a, b, c, M, N, K, - a.stride(0), a.stride(1), - b.stride(0), b.stride(1), - c.stride(0), c.stride(1), - QUANTIZED_M = M // 64, - FLOAT = FLOAT + a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), + QUANTIZED_M = M // 64, FLOAT = FLOAT, EVEN_K = (K % 128 == 0), + ) + return c + +@mm_8bit.register_fake +def _(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + out_dtype = torch.float32 if a.dtype == torch.float8_e4m3fn else torch.int32 + return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype) + + +@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_AUTOTUNE_KEY, cache_results=True) +@triton.jit +def _scaled_mm_kernel( + a_ptr, b_ptr, c_ptr, scale_ptr, + M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, +): + accumulator, offs_cm, offs_cn = _mm_accumulate( + a_ptr, b_ptr, M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, + BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, + FLOAT, EVEN_K, + ) + + #per-row scale on axis 0 (M), fold into the epilogue and cast to the output (compute) dtype directly + scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) + result = accumulator.to(tl.float32) * scale[:, None] + result = result.to(c_ptr.dtype.element_ty) + + _store_c(c_ptr, result, offs_cm, offs_cn, M, N, stride_cm) + +announce_autotuning(_scaled_mm_kernel, name="8-bit scaled matmul") + +@torch.library.custom_op("ot_quant::scaled_mm_8bit", mutates_args=()) +def scaled_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + b, c, scale, M, N, K, FLOAT = _prepare_scaled_mm(a, b, scale, out_dtype) + + def grid(META): + return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) + _scaled_mm_kernel[grid]( + a, b, c, scale, + M, N, K, + a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), + QUANTIZED_M = M // 64, FLOAT = FLOAT, EVEN_K = (K % 128 == 0), ) return c + +@scaled_mm_8bit.register_fake +def _(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype) + + +#_scaled_mm_kernel's epilogue with a second dequant scale on axis 1 (N), for callers whose weight +#is quantized axiswise rather than tensorwise (LinearGGUFA8), where the weight scale is one entry +#per output column and so cannot be folded into the per-row scale +@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_AUTOTUNE_KEY, cache_results=True) +@triton.jit +def _rowcol_scaled_mm_kernel( + a_ptr, b_ptr, c_ptr, scale_ptr, scale_n_ptr, + M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, +): + accumulator, offs_cm, offs_cn = _mm_accumulate( + a_ptr, b_ptr, M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, + BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, + FLOAT, EVEN_K, + ) + + scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) + scale_n = tl.load(scale_n_ptr + offs_cn, mask=offs_cn < N, other=0.0) + result = accumulator.to(tl.float32) * scale[:, None] * scale_n[None, :] + result = result.to(c_ptr.dtype.element_ty) + + _store_c(c_ptr, result, offs_cm, offs_cn, M, N, stride_cm) + +announce_autotuning(_rowcol_scaled_mm_kernel, name="8-bit row/column scaled matmul") + +@torch.library.custom_op("ot_quant::rowcol_scaled_mm_8bit", mutates_args=()) +def rowcol_scaled_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + b, c, scale, scale_n, M, N, K, FLOAT = _prepare_rowcol_scaled_mm(a, b, scale, scale_n, out_dtype) + + def grid(META): + return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) + _rowcol_scaled_mm_kernel[grid]( + a, b, c, scale, scale_n, + M, N, K, + a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), + QUANTIZED_M = M // 64, FLOAT = FLOAT, EVEN_K = (K % 128 == 0), + ) + return c + +@rowcol_scaled_mm_8bit.register_fake +def _(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype) + + +#rank-tile width is min(next_pow2(R), cap): the cap bounds the staged tile, and with it the +#epilogue's shared memory, so occupancy stays rank-independent. 64 keeps +#2*BLOCK_R*(BLOCK_M+BLOCK_N) inside the main mm's pool on Ada; 128 overflowed it (backward OOM) +_LORA_BLOCK_R_CAP = 64 + + +def _prepare_lora(a: torch.Tensor, b: torch.Tensor, xd: torch.Tensor, up: torch.Tensor): + R = xd.shape[1] + assert xd.shape[0] == a.shape[0] and up.shape[0] == R and up.shape[1] == b.shape[1], "Incompatible low-rank dimensions" + assert xd.stride(1) == 1, "xd must be contiguous along the rank axis" + assert xd.dtype == up.dtype + return R, min(max(16, triton.next_power_of_2(R)), _LORA_BLOCK_R_CAP) + +@triton.jit +def _add_lora( + result, xd_ptr, up_ptr, offs_cm, offs_cn, M, N, R, + stride_xdm, stride_upr, stride_upn, + BLOCK_R: tl.constexpr, R_TILES: tl.constexpr, +): + for r0 in tl.static_range(R_TILES): + offs_r = r0 * BLOCK_R + tl.arange(0, BLOCK_R) + xd_ptrs = xd_ptr + offs_cm[:, None] * stride_xdm + offs_r[None, :] + up_ptrs = up_ptr + offs_r[:, None] * stride_upr + offs_cn[None, :] * stride_upn + xd = tl.load(xd_ptrs, mask=(offs_cm[:, None] < M) & (offs_r[None, :] < R), other=0.0) + up = tl.load(up_ptrs, mask=(offs_r[:, None] < R) & (offs_cn[None, :] < N), other=0.0) + result += tl.dot(xd, up) + return result + + +@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_LORA_AUTOTUNE_KEY, cache_results=True) +@triton.jit +def _scaled_lora_mm_kernel( + a_ptr, b_ptr, c_ptr, scale_ptr, xd_ptr, up_ptr, + M, N, K, R, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_xdm, stride_upr, stride_upn, + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, + BLOCK_R: tl.constexpr, R_TILES: tl.constexpr, +): + accumulator, offs_cm, offs_cn = _mm_accumulate( + a_ptr, b_ptr, M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, + BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, + FLOAT, EVEN_K, + ) + + scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) + result = accumulator.to(tl.float32) * scale[:, None] + + result = _add_lora(result, xd_ptr, up_ptr, offs_cm, offs_cn, M, N, R, + stride_xdm, stride_upr, stride_upn, BLOCK_R, R_TILES) + result = result.to(c_ptr.dtype.element_ty) + + _store_c(c_ptr, result, offs_cm, offs_cn, M, N, stride_cm) + +announce_autotuning(_scaled_lora_mm_kernel, name="8-bit scaled LoRA matmul") + +@torch.library.custom_op("ot_quant::scaled_lora_mm_8bit", mutates_args=()) +def scaled_lora_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + #returns (a @ b) * scale + xd @ up in out_dtype (rank-tiled epilogue). xd is (M, r), up is + #(r, N) (the lora_up weight transposed, alpha folded in) + R, block_r = _prepare_lora(a, b, xd, up) + b, c, scale, M, N, K, FLOAT = _prepare_scaled_mm(a, b, scale, out_dtype) + + def grid(META): + return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) + _scaled_lora_mm_kernel[grid]( + a, b, c, scale, xd, up, + M, N, K, R, + a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), xd.stride(0), up.stride(0), up.stride(1), + QUANTIZED_M = M // 64, FLOAT = FLOAT, EVEN_K = (K % 128 == 0), + BLOCK_R = block_r, R_TILES = triton.cdiv(R, block_r), + ) + return c + +@scaled_lora_mm_8bit.register_fake +def _(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype) + + +@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_LORA_AUTOTUNE_KEY, cache_results=True) +@triton.jit +def _rowcol_scaled_lora_mm_kernel( + a_ptr, b_ptr, c_ptr, scale_ptr, scale_n_ptr, xd_ptr, up_ptr, + M, N, K, R, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_xdm, stride_upr, stride_upn, + BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, + BLOCK_R: tl.constexpr, R_TILES: tl.constexpr, +): + accumulator, offs_cm, offs_cn = _mm_accumulate( + a_ptr, b_ptr, M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, + BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, + FLOAT, EVEN_K, + ) + + scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) + scale_n = tl.load(scale_n_ptr + offs_cn, mask=offs_cn < N, other=0.0) + result = accumulator.to(tl.float32) * scale[:, None] * scale_n[None, :] + + result = _add_lora(result, xd_ptr, up_ptr, offs_cm, offs_cn, M, N, R, + stride_xdm, stride_upr, stride_upn, BLOCK_R, R_TILES) + result = result.to(c_ptr.dtype.element_ty) + + _store_c(c_ptr, result, offs_cm, offs_cn, M, N, stride_cm) + +announce_autotuning(_rowcol_scaled_lora_mm_kernel, name="8-bit row/column scaled LoRA matmul") + +@torch.library.custom_op("ot_quant::rowcol_scaled_lora_mm_8bit", mutates_args=()) +def rowcol_scaled_lora_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + R, block_r = _prepare_lora(a, b, xd, up) + b, c, scale, scale_n, M, N, K, FLOAT = _prepare_rowcol_scaled_mm(a, b, scale, scale_n, out_dtype) + + def grid(META): + return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) + _rowcol_scaled_lora_mm_kernel[grid]( + a, b, c, scale, scale_n, xd, up, + M, N, K, R, + a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), xd.stride(0), up.stride(0), up.stride(1), + QUANTIZED_M = M // 64, FLOAT = FLOAT, EVEN_K = (K % 128 == 0), + BLOCK_R = block_r, R_TILES = triton.cdiv(R, block_r), + ) + return c + +@rowcol_scaled_lora_mm_8bit.register_fake +def _(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype) From b75952b2d0a43bf0f675090477afbaca81f84674 Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Sat, 8 Aug 2026 23:56:03 +0200 Subject: [PATCH 3/8] Cut the 8-bit mm autotune key count M is batch*sequence, so unlike N and K it is data-dependent and unbounded: it moves with resolution, frame count, batch size and, on models that prune prompt padding, the longest caption in the batch. Keying the autotuner on M // 64 makes the number of tuning keys grow linearly with M, so a run keeps paying for fresh autotune passes as sequence lengths vary. QUANTIZED_M now uses M.bit_length(), bucketing per doubling instead, which is logarithmic in M. Proportional resolution is the right shape here because the winning config is decided by block count against SM count, which is linear in M -- equal ratios of M matter equally at every scale, while a fixed stride is too fine at large M and too coarse at small. The LoRA epilogue additionally drops stride_upr from its key. That stride does vary -- it is 1 in the forward, where up is a transposed view of lora_up, and N in the backward, where up is lora_down stored row-major -- but keying on it doubles the key count on every model. The two layouts do load differently, and the slab is only BLOCK_R x BLOCK_N and is reused by every M block in the group: too little traffic next to A and B to move the tile choice. Co-Authored-By: Claude Fable 5 --- modules/util/triton_mm_8bit.py | 28 ++++++++++++++++++---------- 1 file changed, 18 insertions(+), 10 deletions(-) diff --git a/modules/util/triton_mm_8bit.py b/modules/util/triton_mm_8bit.py index 62d26dc50..0d2c9af9f 100644 --- a/modules/util/triton_mm_8bit.py +++ b/modules/util/triton_mm_8bit.py @@ -84,16 +84,24 @@ def grid(META): _AUTOTUNE_KEY = [ - 'QUANTIZED_M', #only tune roughly on M, because M is the transformer sequence length - can vary on data + #M is batch*sequence, so unlike N and K it is data-dependent and unbounded: it moves with + #resolution, frame count, batch size and (on models that prune prompt padding) the longest + #caption in the batch. bucketing it per doubling keeps the number of tuning keys logarithmic + #in M instead of linear. proportional resolution is the right shape because the winning config + #is decided by block count against SM count, which is linear in M - so equal ratios of M + #matter equally at every scale, and a fixed stride is too fine at large M and too coarse at small + 'QUANTIZED_M', 'N', 'K', 'stride_bk' #use stride of b as key, to autotune again for a strided rhs matrix (backward pass) ] -#the LoRA epilogue keys on the rank tiling as well, and on stride_upr for the same reason as -#stride_bk above: up arrives row-major (r, N) when the caller's cast to the compute dtype is a -#real copy, and as a transposed view when it is a no-op, and the two layouts load differently -_LORA_AUTOTUNE_KEY = [*_AUTOTUNE_KEY, 'R_TILES', 'BLOCK_R', 'stride_upr'] +#the LoRA epilogue keys on the rank tiling as well. up's row stride is deliberately not a key even +#though it varies: it is 1 in the forward (up is a transposed view of lora_up) and N in the backward +#(up is lora_down, stored row-major), so keying on it would double the key count on every model. +#the two layouts load differently but the slab is only BLOCK_R x BLOCK_N and is reused by every +#M block in the group, too little traffic next to A and B to move the tile choice +_LORA_AUTOTUNE_KEY = [*_AUTOTUNE_KEY, 'R_TILES', 'BLOCK_R'] #configs for the shared _mm_accumulate core: GROUP_SIZE_M is required by its grouped launch #order. Shared memory per config is stages*(BLOCK_M+BLOCK_N)*BLOCK_K bytes and must stay @@ -265,7 +273,7 @@ def grid(META): a, b, c, M, N, K, a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), - QUANTIZED_M = M // 64, FLOAT = FLOAT, EVEN_K = (K % 128 == 0), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), ) return c @@ -310,7 +318,7 @@ def grid(META): a, b, c, scale, M, N, K, a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), - QUANTIZED_M = M // 64, FLOAT = FLOAT, EVEN_K = (K % 128 == 0), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), ) return c @@ -357,7 +365,7 @@ def grid(META): a, b, c, scale, scale_n, M, N, K, a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), - QUANTIZED_M = M // 64, FLOAT = FLOAT, EVEN_K = (K % 128 == 0), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), ) return c @@ -436,7 +444,7 @@ def grid(META): a, b, c, scale, xd, up, M, N, K, R, a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), xd.stride(0), up.stride(0), up.stride(1), - QUANTIZED_M = M // 64, FLOAT = FLOAT, EVEN_K = (K % 128 == 0), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), BLOCK_R = block_r, R_TILES = triton.cdiv(R, block_r), ) return c @@ -486,7 +494,7 @@ def grid(META): a, b, c, scale, scale_n, xd, up, M, N, K, R, a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), xd.stride(0), up.stride(0), up.stride(1), - QUANTIZED_M = M // 64, FLOAT = FLOAT, EVEN_K = (K % 128 == 0), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), BLOCK_R = block_r, R_TILES = triton.cdiv(R, block_r), ) return c From eb85620630a45b5616929ef913b028430bba2459 Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Sun, 9 Aug 2026 18:56:16 +0200 Subject: [PATCH 4/8] Show compile progress in a progress bar instead of a line per compile A cold torch.compile cache announces every frame it compiles, which scrolls the progress bar off the screen. There is no knowable total to build a real progress bar from, so the announcement goes into the postfix of the innermost running bar instead, and is cleared again by that bar's next redraw or by its close. tqdm keeps its bars in an unordered WeakSet, so which of the nested bars is the innermost one cannot be recovered from it. modules/util/tqdm_util.py subclasses tqdm to track that itself and adds show_status() next to tqdm.write(); every tqdm import in the repo now comes from there. Bars owned by mgds are outside this and still draw as before. Also gates the warning filters on OT_DEBUG_WARNINGS, so setting it brings every suppressed message back, and silences the diffusers attention-backend experimental notice. Co-Authored-By: Claude Opus 5 (1M context) --- modules/modelSampler/AnimaSampler.py | 2 +- modules/modelSampler/ChromaSampler.py | 3 +- modules/modelSampler/ErnieSampler.py | 2 +- modules/modelSampler/Flux2Sampler.py | 2 +- modules/modelSampler/FluxSampler.py | 3 +- modules/modelSampler/HiDreamSampler.py | 3 +- modules/modelSampler/HunyuanVideoSampler.py | 2 +- modules/modelSampler/IdeogramSampler.py | 2 +- modules/modelSampler/Krea2Sampler.py | 3 +- modules/modelSampler/PixArtAlphaSampler.py | 3 +- modules/modelSampler/QwenSampler.py | 3 +- modules/modelSampler/SanaSampler.py | 3 +- .../modelSampler/StableDiffusion3Sampler.py | 3 +- .../modelSampler/StableDiffusionSampler.py | 3 +- .../modelSampler/StableDiffusionXLSampler.py | 3 +- modules/modelSampler/WuerstchenSampler.py | 2 +- modules/modelSampler/ZImageSampler.py | 3 +- modules/module/BaseImageCaptionModel.py | 2 +- modules/module/BaseImageMaskModel.py | 2 +- modules/module/GenerateLossesModel.py | 3 +- modules/trainer/BaseTrainer.py | 6 +- modules/trainer/GenericTrainer.py | 3 +- modules/ui/TrainUIController.py | 4 +- modules/util/compile_util.py | 5 +- modules/util/multi_gpu_util.py | 3 +- modules/util/quantization_util.py | 2 +- modules/util/tqdm_util.py | 49 ++++++++++++++++ scripts/util/import_util.py | 57 ++++++++++--------- 28 files changed, 108 insertions(+), 73 deletions(-) create mode 100644 modules/util/tqdm_util.py diff --git a/modules/modelSampler/AnimaSampler.py b/modules/modelSampler/AnimaSampler.py index d5de18ac1..56d5c523e 100644 --- a/modules/modelSampler/AnimaSampler.py +++ b/modules/modelSampler/AnimaSampler.py @@ -12,13 +12,13 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch from diffusers import VaeImageProcessor import numpy as np -from tqdm import tqdm @factory.register(BaseModelSampler, ModelType.ANIMA) diff --git a/modules/modelSampler/ChromaSampler.py b/modules/modelSampler/ChromaSampler.py index 23ef4aee5..b8325d297 100644 --- a/modules/modelSampler/ChromaSampler.py +++ b/modules/modelSampler/ChromaSampler.py @@ -12,11 +12,10 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.CHROMA_1) class ChromaSampler(BaseModelSampler): diff --git a/modules/modelSampler/ErnieSampler.py b/modules/modelSampler/ErnieSampler.py index a9cb57e0a..161390f17 100644 --- a/modules/modelSampler/ErnieSampler.py +++ b/modules/modelSampler/ErnieSampler.py @@ -11,12 +11,12 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch import numpy as np from PIL import Image as PILImage -from tqdm import tqdm @factory.register(BaseModelSampler, ModelType.ERNIE) diff --git a/modules/modelSampler/Flux2Sampler.py b/modules/modelSampler/Flux2Sampler.py index 7ecbd5c83..514e9648c 100644 --- a/modules/modelSampler/Flux2Sampler.py +++ b/modules/modelSampler/Flux2Sampler.py @@ -12,13 +12,13 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch from diffusers.pipelines.flux2.pipeline_flux2 import compute_empirical_mu import numpy as np -from tqdm import tqdm @factory.register(BaseModelSampler, ModelType.FLUX_2) diff --git a/modules/modelSampler/FluxSampler.py b/modules/modelSampler/FluxSampler.py index fc2b2e0c8..8a935be79 100644 --- a/modules/modelSampler/FluxSampler.py +++ b/modules/modelSampler/FluxSampler.py @@ -14,13 +14,12 @@ from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat from modules.util.image_util import load_image +from modules.util.tqdm_util import tqdm import torch from torch import nn from torchvision.transforms import transforms -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.FLUX_DEV_1) @factory.register(BaseModelSampler, ModelType.FLUX_FILL_DEV_1) diff --git a/modules/modelSampler/HiDreamSampler.py b/modules/modelSampler/HiDreamSampler.py index c5b30723b..e04209a00 100644 --- a/modules/modelSampler/HiDreamSampler.py +++ b/modules/modelSampler/HiDreamSampler.py @@ -12,11 +12,10 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.HI_DREAM_FULL) class HiDreamSampler(BaseModelSampler): diff --git a/modules/modelSampler/HunyuanVideoSampler.py b/modules/modelSampler/HunyuanVideoSampler.py index 10b22bfc9..49cb0b426 100644 --- a/modules/modelSampler/HunyuanVideoSampler.py +++ b/modules/modelSampler/HunyuanVideoSampler.py @@ -12,11 +12,11 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch from PIL import Image -from tqdm import tqdm @factory.register(BaseModelSampler, ModelType.HUNYUAN_VIDEO) diff --git a/modules/modelSampler/IdeogramSampler.py b/modules/modelSampler/IdeogramSampler.py index cb253db0e..2672837e7 100644 --- a/modules/modelSampler/IdeogramSampler.py +++ b/modules/modelSampler/IdeogramSampler.py @@ -11,6 +11,7 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch @@ -18,7 +19,6 @@ import numpy as np from PIL import Image as PILImage -from tqdm import tqdm @factory.register(BaseModelSampler, ModelType.IDEOGRAM_4) diff --git a/modules/modelSampler/Krea2Sampler.py b/modules/modelSampler/Krea2Sampler.py index b83205a93..eb0853bb8 100644 --- a/modules/modelSampler/Krea2Sampler.py +++ b/modules/modelSampler/Krea2Sampler.py @@ -13,13 +13,12 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch from diffusers import Krea2Pipeline -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.KREA_2) class Krea2Sampler(BaseModelSampler): diff --git a/modules/modelSampler/PixArtAlphaSampler.py b/modules/modelSampler/PixArtAlphaSampler.py index f8c38f135..58db739ec 100644 --- a/modules/modelSampler/PixArtAlphaSampler.py +++ b/modules/modelSampler/PixArtAlphaSampler.py @@ -11,11 +11,10 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.PIXART_ALPHA) @factory.register(BaseModelSampler, ModelType.PIXART_SIGMA) diff --git a/modules/modelSampler/QwenSampler.py b/modules/modelSampler/QwenSampler.py index c18eece7e..07e66f202 100644 --- a/modules/modelSampler/QwenSampler.py +++ b/modules/modelSampler/QwenSampler.py @@ -13,11 +13,10 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.QWEN) class QwenSampler(BaseModelSampler): diff --git a/modules/modelSampler/SanaSampler.py b/modules/modelSampler/SanaSampler.py index f4089a222..ee7df291d 100644 --- a/modules/modelSampler/SanaSampler.py +++ b/modules/modelSampler/SanaSampler.py @@ -12,11 +12,10 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.SANA) class SanaSampler(BaseModelSampler): diff --git a/modules/modelSampler/StableDiffusion3Sampler.py b/modules/modelSampler/StableDiffusion3Sampler.py index f21f34627..67fef6baa 100644 --- a/modules/modelSampler/StableDiffusion3Sampler.py +++ b/modules/modelSampler/StableDiffusion3Sampler.py @@ -12,11 +12,10 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.STABLE_DIFFUSION_3) @factory.register(BaseModelSampler, ModelType.STABLE_DIFFUSION_35) diff --git a/modules/modelSampler/StableDiffusionSampler.py b/modules/modelSampler/StableDiffusionSampler.py index 791290edc..6f49ef737 100644 --- a/modules/modelSampler/StableDiffusionSampler.py +++ b/modules/modelSampler/StableDiffusionSampler.py @@ -12,13 +12,12 @@ from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat from modules.util.image_util import load_image +from modules.util.tqdm_util import tqdm import torch from torch import nn from torchvision.transforms import transforms -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.STABLE_DIFFUSION_15) @factory.register(BaseModelSampler, ModelType.STABLE_DIFFUSION_15_INPAINTING) diff --git a/modules/modelSampler/StableDiffusionXLSampler.py b/modules/modelSampler/StableDiffusionXLSampler.py index 93d9f23d3..603ef9be6 100644 --- a/modules/modelSampler/StableDiffusionXLSampler.py +++ b/modules/modelSampler/StableDiffusionXLSampler.py @@ -12,13 +12,12 @@ from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat from modules.util.image_util import load_image +from modules.util.tqdm_util import tqdm import torch from torch import nn from torchvision.transforms import transforms -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.STABLE_DIFFUSION_XL_10_BASE) @factory.register(BaseModelSampler, ModelType.STABLE_DIFFUSION_XL_10_BASE_INPAINTING) diff --git a/modules/modelSampler/WuerstchenSampler.py b/modules/modelSampler/WuerstchenSampler.py index a1e9d8ba0..d5c833031 100644 --- a/modules/modelSampler/WuerstchenSampler.py +++ b/modules/modelSampler/WuerstchenSampler.py @@ -11,11 +11,11 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch from PIL import Image -from tqdm import tqdm @factory.register(BaseModelSampler, ModelType.WUERSTCHEN_2) diff --git a/modules/modelSampler/ZImageSampler.py b/modules/modelSampler/ZImageSampler.py index 0e001df2a..371dd20c5 100644 --- a/modules/modelSampler/ZImageSampler.py +++ b/modules/modelSampler/ZImageSampler.py @@ -12,11 +12,10 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - @factory.register(BaseModelSampler, ModelType.Z_IMAGE) class ZImageSampler(BaseModelSampler): diff --git a/modules/module/BaseImageCaptionModel.py b/modules/module/BaseImageCaptionModel.py index 2dfcf4a87..39cf79703 100644 --- a/modules/module/BaseImageCaptionModel.py +++ b/modules/module/BaseImageCaptionModel.py @@ -6,9 +6,9 @@ from modules.util import path_util from modules.util.image_util import load_image +from modules.util.tqdm_util import tqdm from PIL import Image -from tqdm import tqdm class CaptionSample: diff --git a/modules/module/BaseImageMaskModel.py b/modules/module/BaseImageMaskModel.py index 017ec7dd3..f5cfbfabb 100644 --- a/modules/module/BaseImageMaskModel.py +++ b/modules/module/BaseImageMaskModel.py @@ -5,13 +5,13 @@ from modules.util import path_util from modules.util.image_util import load_image +from modules.util.tqdm_util import tqdm import torch from torch import Tensor from torchvision.transforms import transforms from PIL import Image -from tqdm import tqdm class MaskSample: diff --git a/modules/module/GenerateLossesModel.py b/modules/module/GenerateLossesModel.py index d4c821b74..ea90beddc 100644 --- a/modules/module/GenerateLossesModel.py +++ b/modules/module/GenerateLossesModel.py @@ -8,12 +8,11 @@ from modules.util import create from modules.util.config.TrainConfig import QuantizationConfig, TrainConfig from modules.util.torch_util import torch_gc +from modules.util.tqdm_util import tqdm from modules.util.TrainProgress import TrainProgress import torch -from tqdm import tqdm - class GenerateLossesModel: """Based on train args, writes a JSON instead of a model with filenames mapped to losses, diff --git a/modules/trainer/BaseTrainer.py b/modules/trainer/BaseTrainer.py index e4b7a3b29..87a96d7d0 100644 --- a/modules/trainer/BaseTrainer.py +++ b/modules/trainer/BaseTrainer.py @@ -97,10 +97,8 @@ def _start_tensorboard(self): if self.config.tensorboard_expose: tensorboard_args.append("--bind_all") - # Discard the tensorboard child's stdout/stderr: the TF-not-found notice, the - # experimental-data-loading NOTE and the serving banner are all noise, and the - # UI already exposes the tensorboard URL. Popen still raises if the executable - # is missing, so a real launch failure is not hidden. + # discard the child's banner and notices; the UI already shows the tensorboard URL. + # Popen still raises if the executable is missing. self.tensorboard_subprocess = subprocess.Popen( tensorboard_args, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, ) diff --git a/modules/trainer/GenericTrainer.py b/modules/trainer/GenericTrainer.py index 509811d2c..cce3ec51d 100644 --- a/modules/trainer/GenericTrainer.py +++ b/modules/trainer/GenericTrainer.py @@ -32,6 +32,7 @@ from modules.util.profiling_util import PeakMemoryRecorder, TorchMemoryRecorder, TorchProfiler from modules.util.time_util import get_string_timestamp from modules.util.torch_util import torch_gc +from modules.util.tqdm_util import tqdm from modules.util.TrainProgress import TrainProgress import torch @@ -41,8 +42,6 @@ from torch.utils.tensorboard import SummaryWriter from torchvision.transforms.functional import pil_to_tensor -from tqdm import tqdm - class GenericTrainer(BaseTrainer): model_loader: BaseModelLoader diff --git a/modules/ui/TrainUIController.py b/modules/ui/TrainUIController.py index 4fa3fb688..895be42fa 100644 --- a/modules/ui/TrainUIController.py +++ b/modules/ui/TrainUIController.py @@ -108,9 +108,7 @@ def _start_always_on_tensorboard(self): if self.train_config.tensorboard_expose: tensorboard_args.append("--bind_all") - # Discard the tensorboard child's stdout/stderr: the TF-not-found notice, the - # experimental-data-loading NOTE and the serving banner are all noise, and the - # UI already exposes the tensorboard URL. + # discard the child's banner and notices; the UI already shows the tensorboard URL. try: self.always_on_tensorboard_subprocess = subprocess.Popen( tensorboard_args, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, diff --git a/modules/util/compile_util.py b/modules/util/compile_util.py index 0a09f1023..ffada3b2a 100644 --- a/modules/util/compile_util.py +++ b/modules/util/compile_util.py @@ -1,9 +1,10 @@ +from modules.util.tqdm_util import tqdm + import torch import torch._dynamo.callback import torch.utils._sympy.functions from sympy import S -from tqdm import tqdm #code from https://github.com/pytorch/pytorch/blob/ed82d5fcfd80110565f69130f286c7bfec6db2dc/torch/utils/_sympy/functions.py#L481 @@ -95,7 +96,7 @@ def init_compile(): def _on_compile_start(args: "torch._dynamo.callback.CallbackArgs") -> None: frame_id, _, frame_compile_id = args.compile_id.partition("/") direction = "backward" if args.callback_trigger == torch._dynamo.callback.CallbackTrigger.LAZY_BACKWARD else "forward" - tqdm.write(f"[torch.compile] compiling kernel {frame_id} {direction} (variant #{frame_compile_id or 0})...") + tqdm.show_status(f"compiling kernel {frame_id} {direction} (variant #{frame_compile_id or 0})...") torch._dynamo.callback.on_compile_start(_on_compile_start) diff --git a/modules/util/multi_gpu_util.py b/modules/util/multi_gpu_util.py index 962434b65..545a90ef8 100644 --- a/modules/util/multi_gpu_util.py +++ b/modules/util/multi_gpu_util.py @@ -3,11 +3,10 @@ from modules.util.bf16_stochastic_rounding import copy_stochastic_ from modules.util.commands.TrainCommands import TrainCommands from modules.util.enum.GradientReducePrecision import GradientReducePrecision +from modules.util.tqdm_util import tqdm import torch -from tqdm import tqdm - def is_enabled() -> bool: return torch.distributed.is_available() and torch.distributed.is_initialized() diff --git a/modules/util/quantization_util.py b/modules/util/quantization_util.py index d7229687d..d405802b0 100644 --- a/modules/util/quantization_util.py +++ b/modules/util/quantization_util.py @@ -9,6 +9,7 @@ from modules.util.config.TrainConfig import QuantizationConfig, TrainConfig from modules.util.enum.DataType import DataType from modules.util.ModuleFilter import ModuleFilter +from modules.util.tqdm_util import tqdm import torch from torch import Tensor, nn @@ -16,7 +17,6 @@ from diffusers.quantizers.gguf.utils import GGUFLinear, dequantize_gguf_tensor import accelerate -from tqdm import tqdm try: from modules.module.quantized.LinearNf4 import LinearNf4 diff --git a/modules/util/tqdm_util.py b/modules/util/tqdm_util.py new file mode 100644 index 000000000..838e38fee --- /dev/null +++ b/modules/util/tqdm_util.py @@ -0,0 +1,49 @@ +from tqdm import tqdm as _tqdm + +#progress bars created through the tqdm below, innermost last +_bars = [] + + +class tqdm(_tqdm): + _status = None + + @classmethod + def get_lock(cls): + #tqdm caches the terminal write lock on the class that first asks for it, so a subclass + #would get one of its own and stop serializing against bars drawn by tqdm itself. + return _tqdm.get_lock() + + @classmethod + def show_status(cls, message: str): + #status of long-running work - a compile, an autotune sweep - goes into the innermost bar's + #postfix rather than on a line of its own, and stands until the next postfix write replaces it. + bar = next((bar for bar in reversed(_bars) if not bar.disable), None) + if bar is None: + cls.write(message) + else: + bar.set_postfix_str(message) + bar._status = message + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + _bars.append(self) + + def _clear_status(self, refresh=True): + #anything written over the status - the training loop's loss - is left standing. + if self._status is not None and self.postfix == self._status: + self.set_postfix_str("", refresh=refresh) + self._status = None + + def update(self, n=1): + #the status describes work that was running while the bar stood still, so the step that + #follows it is the point where it stops being current. + self._clear_status(refresh=False) + return super().update(n) + + def close(self): + self._clear_status() + super().close() + #compared by identity: tqdm's __eq__ is by screen position, so a bar closed late by __del__ + #would drop whichever live bar has taken over its line. + global _bars + _bars = [bar for bar in _bars if bar is not self] diff --git a/scripts/util/import_util.py b/scripts/util/import_util.py index a0c0a83b6..e269a0505 100644 --- a/scripts/util/import_util.py +++ b/scripts/util/import_util.py @@ -14,37 +14,40 @@ def script_imports(allow_zluda: bool = True): # Silence specific non-actionable startup/compile warnings. A logger filter # targets the exact emitting logger, since a parent logger's filter misses - # records from child loggers. - - # diffusers/transformers chatty logger.warning() lines at import/load time. - logging.getLogger("diffusers.modular_pipelines").addFilter( - lambda record: 'Modular Diffusers is currently an experimental feature' not in record.getMessage() - ) - # The subject of these two is interpolated into the message, so match the whole - # sentence with .* standing in for the runtime value. - logging.getLogger("diffusers.configuration_utils").addFilter( - lambda record: not re.search( - r"The config attributes .* were passed to .*, but are not expected and will be ignored", - record.getMessage(), + # records from child loggers. Set OT_DEBUG_WARNINGS to see them all. + if not os.environ.get("OT_DEBUG_WARNINGS"): + # diffusers/transformers chatty logger.warning() lines at import/load time. + logging.getLogger("diffusers.modular_pipelines").addFilter( + lambda record: 'Modular Diffusers is currently an experimental feature' not in record.getMessage() + ) + # The subject of these two is interpolated into the message, so match the whole + # sentence with .* standing in for the runtime value. + logging.getLogger("diffusers.configuration_utils").addFilter( + lambda record: not re.search( + r"The config attributes .* were passed to .*, but are not expected and will be ignored", + record.getMessage(), + ) ) - ) - logging.getLogger("transformers.modeling_utils").addFilter( - lambda record: not re.search( - r"`loss_type=.*` was set in the config but it is unrecognized", record.getMessage() + logging.getLogger("diffusers.models.modeling_utils").addFilter( + lambda record: 'Attention backends are an experimental feature' not in record.getMessage() + ) + logging.getLogger("transformers.modeling_utils").addFilter( + lambda record: not re.search( + r"`loss_type=.*` was set in the config but it is unrecognized", record.getMessage() + ) ) - ) - # A dependency still calls hf_hub_download with the removed local_dir_use_symlinks - # argument; the deprecation warning is not actionable. - warnings.filterwarnings("ignore", message=r".*local_dir_use_symlinks.*") + # A dependency still calls hf_hub_download with the removed local_dir_use_symlinks + # argument; the deprecation warning is not actionable. + warnings.filterwarnings("ignore", message=r".*local_dir_use_symlinks.*") - # torch.compile emits performance notes when inductor falls back or can't use a - # fast path; harmless and noisy for normal runs. The SMs note is a logger.warning() - # on its exact emitting logger; the complex-operators note is a warnings.warn(). - warnings.filterwarnings("ignore", message=r".*does not support code generation for complex operators.*") - logging.getLogger("torch._inductor.utils").addFilter( - lambda record: 'Not enough SMs to use max_autotune_gemm mode' not in record.getMessage() - ) + # torch.compile emits performance notes when inductor falls back or can't use a + # fast path; harmless and noisy for normal runs. The SMs note is a logger.warning() + # on its exact emitting logger; the complex-operators note is a warnings.warn(). + warnings.filterwarnings("ignore", message=r".*does not support code generation for complex operators.*") + logging.getLogger("torch._inductor.utils").addFilter( + lambda record: 'Not enough SMs to use max_autotune_gemm mode' not in record.getMessage() + ) # Insert ourselves as the highest-priority library path, so our modules are # always found without any risk of being shadowed by another import path. From b091a07f5427816cfdfc7a7593cee599005eb0d8 Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Sun, 9 Aug 2026 19:16:01 +0200 Subject: [PATCH 5/8] Announce autotuning in the progress bar instead of a line per variant A cold autotune cache benchmarks each kernel once per shape key, and every one of those sweeps announced itself on a line of its own. Route them through tqdm.show_status(), so the message sits in the innermost progress bar while the sweep runs and is cleared once the bar moves on. Co-Authored-By: Claude Opus 5 (1M context) --- modules/util/triton_mm_8bit.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/modules/util/triton_mm_8bit.py b/modules/util/triton_mm_8bit.py index 0d2c9af9f..83070f79e 100644 --- a/modules/util/triton_mm_8bit.py +++ b/modules/util/triton_mm_8bit.py @@ -7,22 +7,23 @@ #_scaled_mm_kernel folds a per-row dequant scale in and casts to the output dtype, and #_scaled_lora_mm_kernel additionally fuses a low-rank update into the same tile. +from modules.util.tqdm_util import tqdm + import torch import triton import triton.language as tl -from tqdm import tqdm def announce_autotuning(kernel, name=None): - prefix = f"[triton] Autotuning {name} " if name else "[triton] Autotuning " + prefix = f"autotuning {name} " if name else "autotuning " orig_check_disk_cache = kernel.check_disk_cache variants = 0 def check_disk_cache(tuning_key, configs, bench_fn): def announced_bench(): nonlocal variants variants += 1 - tqdm.write(f"{prefix}variant #{variants}...") + tqdm.show_status(f"{prefix}variant #{variants}...") bench_fn() return orig_check_disk_cache(tuning_key, configs, announced_bench) kernel.check_disk_cache = check_disk_cache From 5188c23e19ce86ed9c3bfb94855bd44ca0fb6308 Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Mon, 10 Aug 2026 23:34:56 +0200 Subject: [PATCH 6/8] Use Blackwell's block-scaled fp8 mma The legacy fp8 mma issues at half rate with an fp32 accumulator on sm_120, while the block-scaled mxf8f6f4 instruction Blackwell added runs at the full 8-bit tensor core rate. Reaching it through tl.dot_scaled with every scale set to the ue8m0 encoding of 1.0 computes exactly the same product, so the results are bit-identical to what the kernels produce today - it is purely an instruction swap. All five entry points share _mm_accumulate, so the swap lands once and all five benefit. Pre-Blackwell has no such instruction and triton emulates tl.dot_scaled with a bf16 mma, which is slower than the plain tl.dot path, so MXFP8_MMA is decided from the compute capability and sm_89 keeps its current path unchanged. Co-Authored-By: Claude Opus 5 (1M context) --- modules/util/triton_mm_8bit.py | 56 +++++++++++++++++++++++----------- 1 file changed, 38 insertions(+), 18 deletions(-) diff --git a/modules/util/triton_mm_8bit.py b/modules/util/triton_mm_8bit.py index 83070f79e..2a8744317 100644 --- a/modules/util/triton_mm_8bit.py +++ b/modules/util/triton_mm_8bit.py @@ -15,6 +15,13 @@ import triton.language as tl +#Blackwell's block-scaled fp8 mma (mxf8f6f4) runs at the full 8-bit tensor core rate, the legacy +#fp8 mma only at half rate. Pre-Blackwell has no such instruction and triton emulates +#tl.dot_scaled with a bf16 mma, which is slower than plain tl.dot - so pick by compute capability +def _prefer_mxfp8(device: torch.device) -> bool: + return torch.cuda.get_device_capability(device)[0] >= 12 + + def announce_autotuning(kernel, name=None): prefix = f"autotuning {name} " if name else "autotuning " orig_check_disk_cache = kernel.check_disk_cache @@ -137,7 +144,7 @@ def _mm_accumulate( M, N, K, stride_am, stride_ak, stride_bk, stride_bn, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, - FLOAT: tl.constexpr, EVEN_K: tl.constexpr, + FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, ): #grouped launch order: consecutive pids walk down GROUP_SIZE_M M-blocks before advancing @@ -168,12 +175,22 @@ def _mm_accumulate( accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32 if FLOAT else tl.int32) + #the mma multiplies each group of 32 elements along K by one ue8m0 scale. ue8m0 is a bare + #exponent with bias 127, so the value 127 means a scale of 1.0 and every element is left + #unchanged - the result is the same as an unscaled fp8 matmul + if MXFP8_MMA: + a_scale = tl.full((BLOCK_SIZE_M, BLOCK_SIZE_K // 32), 127, dtype=tl.uint8) + b_scale = tl.full((BLOCK_SIZE_N, BLOCK_SIZE_K // 32), 127, dtype=tl.uint8) + if EVEN_K: for _k in range(K // BLOCK_SIZE_K): a = tl.load(a_ptrs) b = tl.load(b_ptrs) - accumulator = tl.dot(a, b, accumulator, out_dtype=tl.float32 if FLOAT else tl.int32) + if MXFP8_MMA: + accumulator = tl.dot_scaled(a, a_scale, "e4m3", b, b_scale, "e4m3", acc=accumulator) + else: + accumulator = tl.dot(a, b, accumulator, out_dtype=tl.float32 if FLOAT else tl.int32) a_ptrs += BLOCK_SIZE_K * stride_ak b_ptrs += BLOCK_SIZE_K * stride_bk @@ -182,7 +199,10 @@ def _mm_accumulate( a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k*BLOCK_SIZE_K, other=0.0) b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k*BLOCK_SIZE_K, other=0.0) - accumulator = tl.dot(a, b, accumulator, out_dtype=tl.float32 if FLOAT else tl.int32) + if MXFP8_MMA: + accumulator = tl.dot_scaled(a, a_scale, "e4m3", b, b_scale, "e4m3", acc=accumulator) + else: + accumulator = tl.dot(a, b, accumulator, out_dtype=tl.float32 if FLOAT else tl.int32) a_ptrs += BLOCK_SIZE_K * stride_ak b_ptrs += BLOCK_SIZE_K * stride_bk @@ -246,13 +266,13 @@ def _mm_kernel( M, N, K, stride_am, stride_ak, stride_bk, stride_bn, stride_cm, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, - QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, ): accumulator, offs_cm, offs_cn = _mm_accumulate( a_ptr, b_ptr, M, N, K, stride_am, stride_ak, stride_bk, stride_bn, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, - FLOAT, EVEN_K, + FLOAT, EVEN_K, MXFP8_MMA, ) _store_c(c_ptr, accumulator, offs_cm, offs_cn, M, N, stride_cm) @@ -274,7 +294,7 @@ def grid(META): a, b, c, M, N, K, a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), - QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), ) return c @@ -291,13 +311,13 @@ def _scaled_mm_kernel( M, N, K, stride_am, stride_ak, stride_bk, stride_bn, stride_cm, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, - QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, ): accumulator, offs_cm, offs_cn = _mm_accumulate( a_ptr, b_ptr, M, N, K, stride_am, stride_ak, stride_bk, stride_bn, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, - FLOAT, EVEN_K, + FLOAT, EVEN_K, MXFP8_MMA, ) #per-row scale on axis 0 (M), fold into the epilogue and cast to the output (compute) dtype directly @@ -319,7 +339,7 @@ def grid(META): a, b, c, scale, M, N, K, a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), - QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), ) return c @@ -338,13 +358,13 @@ def _rowcol_scaled_mm_kernel( M, N, K, stride_am, stride_ak, stride_bk, stride_bn, stride_cm, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, - QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, ): accumulator, offs_cm, offs_cn = _mm_accumulate( a_ptr, b_ptr, M, N, K, stride_am, stride_ak, stride_bk, stride_bn, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, - FLOAT, EVEN_K, + FLOAT, EVEN_K, MXFP8_MMA, ) scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) @@ -366,7 +386,7 @@ def grid(META): a, b, c, scale, scale_n, M, N, K, a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), - QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), ) return c @@ -411,14 +431,14 @@ def _scaled_lora_mm_kernel( M, N, K, R, stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_xdm, stride_upr, stride_upn, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, - QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, BLOCK_R: tl.constexpr, R_TILES: tl.constexpr, ): accumulator, offs_cm, offs_cn = _mm_accumulate( a_ptr, b_ptr, M, N, K, stride_am, stride_ak, stride_bk, stride_bn, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, - FLOAT, EVEN_K, + FLOAT, EVEN_K, MXFP8_MMA, ) scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) @@ -445,7 +465,7 @@ def grid(META): a, b, c, scale, xd, up, M, N, K, R, a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), xd.stride(0), up.stride(0), up.stride(1), - QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), BLOCK_R = block_r, R_TILES = triton.cdiv(R, block_r), ) return c @@ -462,14 +482,14 @@ def _rowcol_scaled_lora_mm_kernel( M, N, K, R, stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_xdm, stride_upr, stride_upn, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, - QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, BLOCK_R: tl.constexpr, R_TILES: tl.constexpr, ): accumulator, offs_cm, offs_cn = _mm_accumulate( a_ptr, b_ptr, M, N, K, stride_am, stride_ak, stride_bk, stride_bn, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, - FLOAT, EVEN_K, + FLOAT, EVEN_K, MXFP8_MMA, ) scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) @@ -495,7 +515,7 @@ def grid(META): a, b, c, scale, scale_n, xd, up, M, N, K, R, a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), xd.stride(0), up.stride(0), up.stride(1), - QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), BLOCK_R = block_r, R_TILES = triton.cdiv(R, block_r), ) return c From d416028b3dca851e16a51a18ffa603d0ac800ccd Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Tue, 11 Aug 2026 00:15:31 +0200 Subject: [PATCH 7/8] Only use the block-scaled mma on CUDA torch.cuda also serves ROCm devices, where the device capability is the gfx arch number - RDNA4 reports 12 and would take the tl.dot_scaled path, which triton has to emulate there. Co-Authored-By: Claude Opus 5 (1M context) --- modules/util/triton_mm_8bit.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/modules/util/triton_mm_8bit.py b/modules/util/triton_mm_8bit.py index 2a8744317..53fd0682d 100644 --- a/modules/util/triton_mm_8bit.py +++ b/modules/util/triton_mm_8bit.py @@ -17,9 +17,10 @@ #Blackwell's block-scaled fp8 mma (mxf8f6f4) runs at the full 8-bit tensor core rate, the legacy #fp8 mma only at half rate. Pre-Blackwell has no such instruction and triton emulates -#tl.dot_scaled with a bf16 mma, which is slower than plain tl.dot - so pick by compute capability +#tl.dot_scaled with a bf16 mma, which is slower than plain tl.dot - so pick by compute capability. +#On ROCm the capability is the gfx arch number and RDNA4 reports 12, so require CUDA def _prefer_mxfp8(device: torch.device) -> bool: - return torch.cuda.get_device_capability(device)[0] >= 12 + return torch.version.cuda is not None and torch.cuda.get_device_capability(device)[0] >= 12 def announce_autotuning(kernel, name=None): From bcef2c67a656cee7f11ff1f04707eef1dc0684e8 Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Wed, 19 Aug 2026 20:10:27 +0200 Subject: [PATCH 8/8] Collapse the 8-bit mm variants into one kernel The five entry points (mm_8bit, scaled_mm_8bit, rowcol_scaled_mm_8bit and the two lora variants) were the same _mm_accumulate core with a different epilogue bolted on, each with its own Triton kernel, its own custom op and its own autotune cache. They become one mm_8bit(a, b, out_dtype, scale_m=None, scale_n=None, lora_xd=None, lora_up=None) behind a single ot_quant::mm_8bit op. Triton specializes a None argument as a constexpr, so every combination still compiles to the code a separate kernel per epilogue produced. The main loop no longer carries an EVEN_K variant. _prepare_mm zero-pads K up to _K_ALIGN (the largest BLOCK_SIZE_K in the autotune configs) instead, so every block of the loop is full and the loads need no mask; the padded products are zero. Only layers whose width is not a multiple of that alignment pay for the pad copies. Co-Authored-By: Claude Opus 5 (1M context) --- modules/module/quantized/LinearGGUFA8.py | 19 +- modules/module/quantized/LinearW8A8.py | 26 +- modules/util/mm_8bit.py | 38 +-- modules/util/triton_mm_8bit.py | 377 +++++++---------------- 4 files changed, 137 insertions(+), 323 deletions(-) diff --git a/modules/module/quantized/LinearGGUFA8.py b/modules/module/quantized/LinearGGUFA8.py index d11606d53..27f89c949 100644 --- a/modules/module/quantized/LinearGGUFA8.py +++ b/modules/module/quantized/LinearGGUFA8.py @@ -1,6 +1,5 @@ from modules.module.quantized.mixin.LoRAFusableLinearMixin import LoRAFusableLinearMixin from modules.util.mm_8bit import mm_8bit as mm_8bit -from modules.util.mm_8bit import rowcol_scaled_lora_mm_8bit, rowcol_scaled_mm_8bit from modules.util.quantization_util import ( quantize_axiswise, quantize_fp8_axiswise, @@ -42,21 +41,13 @@ def forward_axiswise_postscaled_torch(dtype: torch.dtype, x: Tensor, weight: Ten res_scaled.add_(bias) return res_scaled -@torch.no_grad() -def backward_axiswise_postscaled_triton(dtype: torch.dtype, output: Tensor, weight: Tensor) -> Tensor: - output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=dtype) - w_8, w_scale = quantize_axiswise(weight, dim=0, dtype=dtype) - mm_res = mm_8bit(output_8.contiguous(), w_8) - return mm_res.float().mul_(w_scale).mul_(output_scale).to(output.dtype) - - @torch.no_grad() def forward_axiswise_epiloguescaled_triton(dtype: torch.dtype, x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: x_8, x_scale = quantize_axiswise(x, dim=-1, dtype=dtype) w_8, w_scale = quantize_axiswise(weight, dim=-1, dtype=dtype) - #rowcol_scaled_mm_8bit folds the per-token scale (axis 0) and the per-channel weight scale + #the mm folds the per-token scale (axis 0) and the per-channel weight scale #(axis 1) into the epilogue and returns compute_dtype directly - res_scaled = rowcol_scaled_mm_8bit(x_8, w_8.T, x_scale, w_scale, compute_dtype) + res_scaled = mm_8bit(x_8, w_8.T, out_dtype=compute_dtype, scale_m=x_scale, scale_n=w_scale) if bias is not None: res_scaled.add_(bias) return res_scaled @@ -65,14 +56,14 @@ def forward_axiswise_epiloguescaled_triton(dtype: torch.dtype, x: Tensor, weight def backward_axiswise_epiloguescaled_triton(dtype: torch.dtype, output: Tensor, weight: Tensor) -> Tensor: output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=dtype) w_8, w_scale = quantize_axiswise(weight, dim=0, dtype=dtype) - return rowcol_scaled_mm_8bit(output_8.contiguous(), w_8, output_scale, w_scale, output.dtype) + return mm_8bit(output_8.contiguous(), w_8, out_dtype=output.dtype, scale_m=output_scale, scale_n=w_scale) @torch.no_grad() def forward_axiswise_lora_epiloguescaled_triton(dtype: torch.dtype, x: Tensor, weight: Tensor, bias: Tensor | None, compute_dtype: torch.dtype, x_down: Tensor, lora_up: Tensor) -> Tensor: x_8, x_scale = quantize_axiswise(x, dim=-1, dtype=dtype) w_8, w_scale = quantize_axiswise(weight, dim=-1, dtype=dtype) - res_scaled = rowcol_scaled_lora_mm_8bit(x_8, w_8.T, x_scale, w_scale, x_down, lora_up, compute_dtype) + res_scaled = mm_8bit(x_8, w_8.T, out_dtype=compute_dtype, scale_m=x_scale, scale_n=w_scale, lora_xd=x_down, lora_up=lora_up) if bias is not None: res_scaled.add_(bias) return res_scaled @@ -81,7 +72,7 @@ def forward_axiswise_lora_epiloguescaled_triton(dtype: torch.dtype, x: Tensor, w def backward_axiswise_lora_epiloguescaled_triton(dtype: torch.dtype, output: Tensor, weight: Tensor, grad_x_down_pre: Tensor, lora_down: Tensor) -> Tensor: output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=dtype) w_8, w_scale = quantize_axiswise(weight, dim=0, dtype=dtype) - return rowcol_scaled_lora_mm_8bit(output_8.contiguous(), w_8, output_scale, w_scale, grad_x_down_pre, lora_down.to(grad_x_down_pre.dtype), output.dtype) + return mm_8bit(output_8.contiguous(), w_8, out_dtype=output.dtype, scale_m=output_scale, scale_n=w_scale, lora_xd=grad_x_down_pre, lora_up=lora_down.to(grad_x_down_pre.dtype)) forward_axiswise = forward_axiswise_epiloguescaled_triton diff --git a/modules/module/quantized/LinearW8A8.py b/modules/module/quantized/LinearW8A8.py index a49b7c9f7..36180d5e6 100644 --- a/modules/module/quantized/LinearW8A8.py +++ b/modules/module/quantized/LinearW8A8.py @@ -3,7 +3,6 @@ from modules.module.quantized.mixin.QuantizedLinearMixin import QuantizedLinearMixin from modules.module.quantized.mixin.QuantizedModuleMixin import QuantizedModuleMixin from modules.util.mm_8bit import mm_8bit as mm_8bit -from modules.util.mm_8bit import scaled_lora_mm_8bit, scaled_mm_8bit from modules.util.quantization_util import ( dequantize, quantize_axiswise, @@ -35,22 +34,16 @@ def forward_tokenwise_postscaled_torch(x: Tensor, weight: Tensor, weight_scale: @torch.no_grad() def forward_tokenwise_epiloguescaled_triton(x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype) -> Tensor: x_8, x_scale = quantize_axiswise(x, dim=-1, dtype=weight.dtype) - res_scaled = scaled_mm_8bit(x_8, weight.T, weight_scale * x_scale, compute_dtype) + res_scaled = mm_8bit(x_8, weight.T, out_dtype=compute_dtype, scale_m=weight_scale * x_scale) if bias is not None: res_scaled.add_(bias) return res_scaled -@torch.no_grad() -def backward_tokenwise_postscaled_triton(output: Tensor, weight: Tensor, weight_scale: Tensor) -> Tensor: - output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=weight.dtype) - #almost always, grad outputs are already contiguous and this is a no-op. But there are some grad outputs from SDXL that are non-contiguous: - mm_res = mm_8bit(output_8.contiguous(), weight) - return mm_res.float().mul_(weight_scale * output_scale).to(output.dtype) - @torch.no_grad() def backward_tokenwise_epiloguescaled_triton(output: Tensor, weight: Tensor, weight_scale: Tensor) -> Tensor: output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=weight.dtype) - return scaled_mm_8bit(output_8.contiguous(), weight, weight_scale * output_scale, output.dtype) + #almost always, grad outputs are already contiguous and this is a no-op. But there are some grad outputs from SDXL that are non-contiguous: + return mm_8bit(output_8.contiguous(), weight, out_dtype=output.dtype, scale_m=weight_scale * output_scale) forward_tokenwise = forward_tokenwise_epiloguescaled_triton @@ -60,7 +53,7 @@ def backward_tokenwise_epiloguescaled_triton(output: Tensor, weight: Tensor, wei @torch.no_grad() def forward_tokenwise_lora_epiloguescaled_triton(x: Tensor, weight: Tensor, weight_scale: Tensor, bias: Tensor | None, compute_dtype: torch.dtype, x_down: Tensor, lora_up: Tensor) -> Tensor: x_8, x_scale = quantize_axiswise(x, dim=-1, dtype=weight.dtype) - res_scaled = scaled_lora_mm_8bit(x_8, weight.T, weight_scale * x_scale, x_down, lora_up, compute_dtype) + res_scaled = mm_8bit(x_8, weight.T, out_dtype=compute_dtype, scale_m=weight_scale * x_scale, lora_xd=x_down, lora_up=lora_up) if bias is not None: res_scaled.add_(bias) return res_scaled @@ -68,7 +61,7 @@ def forward_tokenwise_lora_epiloguescaled_triton(x: Tensor, weight: Tensor, weig @torch.no_grad() def backward_tokenwise_lora_epiloguescaled_triton(output: Tensor, weight: Tensor, weight_scale: Tensor, grad_x_down_pre: Tensor, lora_down: Tensor) -> Tensor: output_8, output_scale = quantize_axiswise(output, dim=-1, dtype=weight.dtype) - return scaled_lora_mm_8bit(output_8.contiguous(), weight, weight_scale * output_scale, grad_x_down_pre, lora_down.to(grad_x_down_pre.dtype), output.dtype) + return mm_8bit(output_8.contiguous(), weight, out_dtype=output.dtype, scale_m=weight_scale * output_scale, lora_xd=grad_x_down_pre, lora_up=lora_down.to(grad_x_down_pre.dtype)) forward_tokenwise_lora = forward_tokenwise_lora_epiloguescaled_triton @@ -246,14 +239,13 @@ def benchmark_int8(m, k, n, device = 'cuda', steps = 10000): run_benchmark(lambda: torch._int_mm(x_8, w_8.T), "torch mm int", steps=steps) - run_benchmark(lambda: mm_8bit(x_8, w_8.T), "triton mm int", steps=steps) + run_benchmark(lambda: mm_8bit(x_8, w_8.T, out_dtype=torch.int32), "triton mm int", steps=steps) def torch_backward(a, b): torch._int_mm(a, b.T.contiguous().T) run_benchmark(lambda: torch_backward(y_8, w_8), "torch mm backward int8", steps=steps) - run_benchmark(lambda: mm_8bit(y_8, w_8), "triton mm backward int8", steps=steps) + run_benchmark(lambda: mm_8bit(y_8, w_8, out_dtype=torch.int32), "triton mm backward int8", steps=steps) run_benchmark(lambda: forward_tokenwise_postscaled_torch(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "torch forward int", steps=steps, compile=True) - run_benchmark(lambda: backward_tokenwise_postscaled_triton(y, w_8, w_scale), "triton backward int", steps=steps, compile=True) run_benchmark(lambda: forward_tokenwise_epiloguescaled_triton(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "triton scaled forward int", steps=steps, compile=True) run_benchmark(lambda: backward_tokenwise_epiloguescaled_triton(y, w_8, w_scale), "triton scaled backward int", steps=steps, compile=True) @@ -269,11 +261,11 @@ def benchmark_fp8(m, k, n, device = 'cuda', steps = 10000): one_scale = torch.ones(1, device=device) run_benchmark(lambda: torch._scaled_mm(x_8, w_8.T, out_dtype=torch.bfloat16, scale_a=one_scale.float(), scale_b=w_scale.float()), "torch mm fp8", steps=steps) - run_benchmark(lambda: mm_8bit(x_8, w_8.T), "triton mm fp8", steps=steps) + run_benchmark(lambda: mm_8bit(x_8, w_8.T, out_dtype=torch.float32), "triton mm fp8", steps=steps) def torch_backward(a, b): torch._scaled_mm(a, b.T.contiguous().T, out_dtype=torch.bfloat16, scale_a=one_scale.float(), scale_b=w_scale.float()) run_benchmark(lambda: torch_backward(y_8, w_8), "torch mm backward fp8", steps=steps) - run_benchmark(lambda: mm_8bit(y_8, w_8), "triton mm backward fp8", steps=steps) + run_benchmark(lambda: mm_8bit(y_8, w_8, out_dtype=torch.float32), "triton mm backward fp8", steps=steps) run_benchmark(lambda: forward_tokenwise_postscaled_torch(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "torch forward fp8", steps=steps, compile=True) run_benchmark(lambda: forward_tokenwise_epiloguescaled_triton(x, w_8, w_scale, bias=None, compute_dtype=torch.bfloat16), "triton scaled forward fp8", steps=steps, compile=True) diff --git a/modules/util/mm_8bit.py b/modules/util/mm_8bit.py index e76a2155e..f13452191 100644 --- a/modules/util/mm_8bit.py +++ b/modules/util/mm_8bit.py @@ -1,35 +1,29 @@ import torch try: - from modules.util.triton_mm_8bit import ( - mm_8bit, - rowcol_scaled_lora_mm_8bit, - rowcol_scaled_mm_8bit, - scaled_lora_mm_8bit, - scaled_mm_8bit, - ) + from modules.util.triton_mm_8bit import mm_8bit except ImportError as e: print(str(e) + ", continuing without triton") - def mm_8bit(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + def mm_8bit(a: torch.Tensor, b: torch.Tensor, out_dtype: torch.dtype, + scale_m: torch.Tensor | None = None, scale_n: torch.Tensor | None = None, + lora_xd: torch.Tensor | None = None, lora_up: torch.Tensor | None = None) -> torch.Tensor: assert a.shape[1] == b.shape[0], "Incompatible dimensions" assert a.is_contiguous(), "Matrix A must be contiguous" assert a.dtype == b.dtype, "Incompatible dtypes" assert a.dtype in [torch.int8, torch.float8_e4m3fn] if a.dtype == torch.int8: - return torch._int_mm(a, b) + #cublas rejects an n-major b (the backward's weight) outright, so make it k-major + res = torch._int_mm(a, b.T.contiguous().T) else: one = torch.ones(1, device=a.device) - #out_dtype defaults to a's dtype, which would round the accumulator back to fp8 - return torch._scaled_mm(a, b.T.contiguous().T, scale_a=one, scale_b=one, out_dtype=torch.float32) + res = torch._scaled_mm(a, b.T.contiguous().T, scale_a=one, scale_b=one, out_dtype=torch.float32) - def scaled_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: - return mm_8bit(a, b).float().mul_(scale.reshape(-1, 1)).to(out_dtype) - - def scaled_lora_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: - return mm_8bit(a, b).float().mul_(scale.reshape(-1, 1)).add_(xd @ up).to(out_dtype) - - def rowcol_scaled_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: - return mm_8bit(a, b).float().mul_(scale.reshape(-1, 1)).mul_(scale_n.reshape(1, -1)).to(out_dtype) - - def rowcol_scaled_lora_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: - return mm_8bit(a, b).float().mul_(scale.reshape(-1, 1)).mul_(scale_n.reshape(1, -1)).add_(xd @ up).to(out_dtype) + if scale_m is not None or scale_n is not None or lora_xd is not None: + res = res.float() + if scale_m is not None: + res = res.mul_(scale_m.reshape(-1, 1)) + if scale_n is not None: + res = res.mul_(scale_n.reshape(1, -1)) + if lora_xd is not None: + res = res.add_(lora_xd @ lora_up) + return res.to(out_dtype) diff --git a/modules/util/triton_mm_8bit.py b/modules/util/triton_mm_8bit.py index 53fd0682d..97016205e 100644 --- a/modules/util/triton_mm_8bit.py +++ b/modules/util/triton_mm_8bit.py @@ -1,11 +1,11 @@ #8bit matmul kernels adapted from the Triton tutorial here: #https://triton-lang.org/main/getting-started/tutorials/03-matrix-multiplication.html -#All three mm entry points share the _mm_accumulate compute core (grouped launch order -#for L2 reuse, compile-time layout/divisibility specialization, EVEN_K loop selection) -#and differ only in the epilogue: _mm_kernel stores the raw int32/fp32 product, -#_scaled_mm_kernel folds a per-row dequant scale in and casts to the output dtype, and -#_scaled_lora_mm_kernel additionally fuses a low-rank update into the same tile. +#There is one mm entry point, built on the _mm_accumulate compute core (grouped launch order +#for L2 reuse, compile-time layout/divisibility specialization). +#Everything the callers vary is an optional epilogue argument defaulting to None: a per-row +#dequant scale, a per-column one, a fused low-rank update. Triton specializes a None argument +#as a constexpr, so each combination compiles to what a separate kernel per epilogue would. from modules.util.tqdm_util import tqdm @@ -102,16 +102,11 @@ def grid(META): 'QUANTIZED_M', 'N', 'K', - 'stride_bk' #use stride of b as key, to autotune again for a strided rhs matrix (backward pass) + 'stride_bk', #use stride of b as key, to autotune again for a strided rhs matrix (backward pass) + 'R_TILES', + 'BLOCK_R', ] -#the LoRA epilogue keys on the rank tiling as well. up's row stride is deliberately not a key even -#though it varies: it is 1 in the forward (up is a transposed view of lora_up) and N in the backward -#(up is lora_down, stored row-major), so keying on it would double the key count on every model. -#the two layouts load differently but the slab is only BLOCK_R x BLOCK_N and is reused by every -#M block in the group, too little traffic next to A and B to move the tile choice -_LORA_AUTOTUNE_KEY = [*_AUTOTUNE_KEY, 'R_TILES', 'BLOCK_R'] - #configs for the shared _mm_accumulate core: GROUP_SIZE_M is required by its grouped launch #order. Shared memory per config is stages*(BLOCK_M+BLOCK_N)*BLOCK_K bytes and must stay #under the ~99KB per-CTA limit (identical on sm86/sm89/sm120, i.e. consumer/workstation @@ -137,7 +132,7 @@ def grid(META): #shared compute core of the 8-bit mm kernels: grouped launch order, divisibility hints and -#the EVEN_K-specialized main loop. Returns the raw accumulator plus the output tile offsets; +#the main loop. Returns the raw accumulator plus the output tile offsets; #each entry kernel below adds its own epilogue and stores. @triton.jit def _mm_accumulate( @@ -145,7 +140,7 @@ def _mm_accumulate( M, N, K, stride_am, stride_ak, stride_bk, stride_bn, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, - FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, + FLOAT: tl.constexpr, MXFP8_MMA: tl.constexpr, ): #grouped launch order: consecutive pids walk down GROUP_SIZE_M M-blocks before advancing @@ -183,30 +178,19 @@ def _mm_accumulate( a_scale = tl.full((BLOCK_SIZE_M, BLOCK_SIZE_K // 32), 127, dtype=tl.uint8) b_scale = tl.full((BLOCK_SIZE_N, BLOCK_SIZE_K // 32), 127, dtype=tl.uint8) - if EVEN_K: - for _k in range(K // BLOCK_SIZE_K): - a = tl.load(a_ptrs) - b = tl.load(b_ptrs) - - if MXFP8_MMA: - accumulator = tl.dot_scaled(a, a_scale, "e4m3", b, b_scale, "e4m3", acc=accumulator) - else: - accumulator = tl.dot(a, b, accumulator, out_dtype=tl.float32 if FLOAT else tl.int32) + #K is a multiple of every BLOCK_SIZE_K (see _K_ALIGN), so every block is full and the loads + #need no mask + for _k in range(K // BLOCK_SIZE_K): + a = tl.load(a_ptrs) + b = tl.load(b_ptrs) - a_ptrs += BLOCK_SIZE_K * stride_ak - b_ptrs += BLOCK_SIZE_K * stride_bk - else: - for k in range(tl.cdiv(K, BLOCK_SIZE_K)): - a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k*BLOCK_SIZE_K, other=0.0) - b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k*BLOCK_SIZE_K, other=0.0) + if MXFP8_MMA: + accumulator = tl.dot_scaled(a, a_scale, "e4m3", b, b_scale, "e4m3", acc=accumulator) + else: + accumulator = tl.dot(a, b, accumulator, out_dtype=tl.float32 if FLOAT else tl.int32) - if MXFP8_MMA: - accumulator = tl.dot_scaled(a, a_scale, "e4m3", b, b_scale, "e4m3", acc=accumulator) - else: - accumulator = tl.dot(a, b, accumulator, out_dtype=tl.float32 if FLOAT else tl.int32) - - a_ptrs += BLOCK_SIZE_K * stride_ak - b_ptrs += BLOCK_SIZE_K * stride_bk + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) @@ -223,6 +207,11 @@ def _store_c(c_ptr, value, offs_cm, offs_cn, M, N, stride_cm): tl.store(c_ptrs, value, mask=c_mask) +#the main loop reads whole K blocks, so _prepare_mm zero-pads K up to a multiple of every +#BLOCK_SIZE_K +_K_ALIGN = max(c.kwargs['BLOCK_SIZE_K'] for c in _AUTOTUNE_CONFIGS) +assert triton.next_power_of_2(_K_ALIGN) == _K_ALIGN + def _prepare_mm(a: torch.Tensor, b: torch.Tensor, out_dtype: torch.dtype): assert a.shape[1] == b.shape[0], "Incompatible dimensions" assert a.is_contiguous(), "Matrix A must be contiguous" @@ -237,163 +226,29 @@ def _prepare_mm(a: torch.Tensor, b: torch.Tensor, out_dtype: torch.dtype): #the kernel handles exactly two B layouts: k-major (forward, weight.T) and n-major (backward, weight) B_K_MAJOR = (b.stride(0) == 1) assert B_K_MAJOR or b.stride(1) == 1, "Matrix B must be contiguous along one axis" + + #zero-padding K keeps every block of the main loop full; the padded products are zero. Only + #layers with an odd width pay for the copies (SDXL's 320-wide ones, Sana's 2240) + if K % _K_ALIGN != 0: + pad = _K_ALIGN - K % _K_ALIGN + a = torch.nn.functional.pad(a, (0, pad)) + #padding the storage of b keeps its layout, padding a transposed view would not + b = torch.nn.functional.pad(b.t(), (0, pad)).t() if B_K_MAJOR else torch.nn.functional.pad(b, (0, 0, 0, pad)) + K += pad #n-major B runs the mm ~30% slower; for large M a transpose copy pays for itself (see transpose_8bit) if not B_K_MAJOR and M >= _TRANSPOSE_MIN_M: b = transpose_8bit(b).t() B_K_MAJOR = True c = torch.empty((M, N), device=a.device, dtype=out_dtype) - return b, c, M, N, K, FLOAT + return a, b, c, M, N, K, FLOAT -def _prepare_scaled_mm(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, out_dtype: torch.dtype): - #_prepare_mm plus the per-row scale reshape (the scaled kernels fold it into the epilogue) - b, c, M, N, K, FLOAT = _prepare_mm(a, b, out_dtype) +def _prepare_scale(scale: torch.Tensor, entries: int, axis: str): + #a dequant scale is one fp32 value per row of a or per column of b, read straight from the + #epilogue, so any other shape or dtype is materialized here rather than in the kernel scale = scale.reshape(-1).to(torch.float32).contiguous() - assert scale.shape[0] == M, "scale must have one entry per row of a" - return b, c, scale, M, N, K, FLOAT - - -def _prepare_rowcol_scaled_mm(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, out_dtype: torch.dtype): - b, c, scale, M, N, K, FLOAT = _prepare_scaled_mm(a, b, scale, out_dtype) - scale_n = scale_n.reshape(-1).to(torch.float32).contiguous() - assert scale_n.shape[0] == N, "scale_n must have one entry per column of b" - return b, c, scale, scale_n, M, N, K, FLOAT - - -@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_AUTOTUNE_KEY, cache_results=True) -@triton.jit -def _mm_kernel( - a_ptr, b_ptr, c_ptr, - M, N, K, - stride_am, stride_ak, stride_bk, stride_bn, stride_cm, - BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, - QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, -): - accumulator, offs_cm, offs_cn = _mm_accumulate( - a_ptr, b_ptr, M, N, K, - stride_am, stride_ak, stride_bk, stride_bn, - BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, - FLOAT, EVEN_K, MXFP8_MMA, - ) - - _store_c(c_ptr, accumulator, offs_cm, offs_cn, M, N, stride_cm) - -announce_autotuning(_mm_kernel, name="8-bit matmul") - -#Opaque custom ops work around pytorch#164124: torch.compile otherwise absorbs the traced -#@triton.autotune kernels and freezes the config benchmarked for the first shape. Kept opaque, -#these bodies run eagerly, so Triton's autotuner selects per key and the JIT sees real sizes -@torch.library.custom_op("ot_quant::mm_8bit", mutates_args=()) -def mm_8bit(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: - out_dtype = torch.float32 if a.dtype == torch.float8_e4m3fn else torch.int32 - b, c, M, N, K, FLOAT = _prepare_mm(a, b, out_dtype) - - #1D grid: the kernel derives pid_m/pid_n itself in grouped order for L2 reuse - def grid(META): - return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) - _mm_kernel[grid]( - a, b, c, - M, N, K, - a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), - QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), - ) - return c - -@mm_8bit.register_fake -def _(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: - out_dtype = torch.float32 if a.dtype == torch.float8_e4m3fn else torch.int32 - return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype) - - -@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_AUTOTUNE_KEY, cache_results=True) -@triton.jit -def _scaled_mm_kernel( - a_ptr, b_ptr, c_ptr, scale_ptr, - M, N, K, - stride_am, stride_ak, stride_bk, stride_bn, stride_cm, - BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, - QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, -): - accumulator, offs_cm, offs_cn = _mm_accumulate( - a_ptr, b_ptr, M, N, K, - stride_am, stride_ak, stride_bk, stride_bn, - BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, - FLOAT, EVEN_K, MXFP8_MMA, - ) - - #per-row scale on axis 0 (M), fold into the epilogue and cast to the output (compute) dtype directly - scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) - result = accumulator.to(tl.float32) * scale[:, None] - result = result.to(c_ptr.dtype.element_ty) - - _store_c(c_ptr, result, offs_cm, offs_cn, M, N, stride_cm) - -announce_autotuning(_scaled_mm_kernel, name="8-bit scaled matmul") - -@torch.library.custom_op("ot_quant::scaled_mm_8bit", mutates_args=()) -def scaled_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: - b, c, scale, M, N, K, FLOAT = _prepare_scaled_mm(a, b, scale, out_dtype) - - def grid(META): - return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) - _scaled_mm_kernel[grid]( - a, b, c, scale, - M, N, K, - a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), - QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), - ) - return c - -@scaled_mm_8bit.register_fake -def _(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: - return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype) - - -#_scaled_mm_kernel's epilogue with a second dequant scale on axis 1 (N), for callers whose weight -#is quantized axiswise rather than tensorwise (LinearGGUFA8), where the weight scale is one entry -#per output column and so cannot be folded into the per-row scale -@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_AUTOTUNE_KEY, cache_results=True) -@triton.jit -def _rowcol_scaled_mm_kernel( - a_ptr, b_ptr, c_ptr, scale_ptr, scale_n_ptr, - M, N, K, - stride_am, stride_ak, stride_bk, stride_bn, stride_cm, - BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, - QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, -): - accumulator, offs_cm, offs_cn = _mm_accumulate( - a_ptr, b_ptr, M, N, K, - stride_am, stride_ak, stride_bk, stride_bn, - BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, - FLOAT, EVEN_K, MXFP8_MMA, - ) - - scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) - scale_n = tl.load(scale_n_ptr + offs_cn, mask=offs_cn < N, other=0.0) - result = accumulator.to(tl.float32) * scale[:, None] * scale_n[None, :] - result = result.to(c_ptr.dtype.element_ty) - - _store_c(c_ptr, result, offs_cm, offs_cn, M, N, stride_cm) - -announce_autotuning(_rowcol_scaled_mm_kernel, name="8-bit row/column scaled matmul") - -@torch.library.custom_op("ot_quant::rowcol_scaled_mm_8bit", mutates_args=()) -def rowcol_scaled_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: - b, c, scale, scale_n, M, N, K, FLOAT = _prepare_rowcol_scaled_mm(a, b, scale, scale_n, out_dtype) - - def grid(META): - return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) - _rowcol_scaled_mm_kernel[grid]( - a, b, c, scale, scale_n, - M, N, K, - a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), - QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), - ) - return c - -@rowcol_scaled_mm_8bit.register_fake -def _(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: - return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype) + assert scale.shape[0] == entries, f"scale must have one entry per {axis}" + return scale #rank-tile width is min(next_pow2(R), cap): the cap bounds the staged tile, and with it the @@ -402,12 +257,13 @@ def _(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tens _LORA_BLOCK_R_CAP = 64 -def _prepare_lora(a: torch.Tensor, b: torch.Tensor, xd: torch.Tensor, up: torch.Tensor): - R = xd.shape[1] - assert xd.shape[0] == a.shape[0] and up.shape[0] == R and up.shape[1] == b.shape[1], "Incompatible low-rank dimensions" - assert xd.stride(1) == 1, "xd must be contiguous along the rank axis" - assert xd.dtype == up.dtype - return R, min(max(16, triton.next_power_of_2(R)), _LORA_BLOCK_R_CAP) +def _prepare_lora(lora_xd: torch.Tensor, lora_up: torch.Tensor, M: int, N: int): + R = lora_xd.shape[1] + assert lora_xd.shape[0] == M and lora_up.shape[0] == R and lora_up.shape[1] == N, "Incompatible low-rank dimensions" + assert lora_xd.stride(1) == 1, "lora_xd must be contiguous along the rank axis" + assert lora_xd.dtype == lora_up.dtype + block_r = min(max(16, triton.next_power_of_2(R)), _LORA_BLOCK_R_CAP) + return R, block_r, triton.cdiv(R, block_r), lora_xd.stride(0), lora_up.stride(0), lora_up.stride(1) @triton.jit def _add_lora( @@ -425,102 +281,83 @@ def _add_lora( return result -@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_LORA_AUTOTUNE_KEY, cache_results=True) +@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_AUTOTUNE_KEY, cache_results=True) @triton.jit -def _scaled_lora_mm_kernel( - a_ptr, b_ptr, c_ptr, scale_ptr, xd_ptr, up_ptr, - M, N, K, R, - stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_xdm, stride_upr, stride_upn, +def _mm_kernel( + a_ptr, b_ptr, c_ptr, + M, N, K, + stride_am, stride_ak, stride_bk, stride_bn, stride_cm, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, - QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, - BLOCK_R: tl.constexpr, R_TILES: tl.constexpr, + QUANTIZED_M, FLOAT: tl.constexpr, MXFP8_MMA: tl.constexpr, + scale_m_ptr=None, scale_n_ptr=None, + xd_ptr=None, up_ptr=None, R=None, stride_xdm=None, stride_upr=None, stride_upn=None, + BLOCK_R: tl.constexpr = None, R_TILES: tl.constexpr = None, ): - accumulator, offs_cm, offs_cn = _mm_accumulate( + result, offs_cm, offs_cn = _mm_accumulate( a_ptr, b_ptr, M, N, K, stride_am, stride_ak, stride_bk, stride_bn, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, - FLOAT, EVEN_K, MXFP8_MMA, + FLOAT, MXFP8_MMA, ) - scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) - result = accumulator.to(tl.float32) * scale[:, None] + if scale_m_ptr is not None or scale_n_ptr is not None or xd_ptr is not None: + result = result.to(tl.float32) - result = _add_lora(result, xd_ptr, up_ptr, offs_cm, offs_cn, M, N, R, - stride_xdm, stride_upr, stride_upn, BLOCK_R, R_TILES) - result = result.to(c_ptr.dtype.element_ty) + if scale_m_ptr is not None: + scale_m = tl.load(scale_m_ptr + offs_cm, mask=offs_cm < M, other=0.0) + result = result * scale_m[:, None] - _store_c(c_ptr, result, offs_cm, offs_cn, M, N, stride_cm) + if scale_n_ptr is not None: + scale_n = tl.load(scale_n_ptr + offs_cn, mask=offs_cn < N, other=0.0) + result = result * scale_n[None, :] -announce_autotuning(_scaled_lora_mm_kernel, name="8-bit scaled LoRA matmul") + if xd_ptr is not None: + result = _add_lora(result, xd_ptr, up_ptr, offs_cm, offs_cn, M, N, R, + stride_xdm, stride_upr, stride_upn, BLOCK_R, R_TILES) -@torch.library.custom_op("ot_quant::scaled_lora_mm_8bit", mutates_args=()) -def scaled_lora_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: - #returns (a @ b) * scale + xd @ up in out_dtype (rank-tiled epilogue). xd is (M, r), up is - #(r, N) (the lora_up weight transposed, alpha folded in) - R, block_r = _prepare_lora(a, b, xd, up) - b, c, scale, M, N, K, FLOAT = _prepare_scaled_mm(a, b, scale, out_dtype) + _store_c(c_ptr, result.to(c_ptr.dtype.element_ty), offs_cm, offs_cn, M, N, stride_cm) - def grid(META): - return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) - _scaled_lora_mm_kernel[grid]( - a, b, c, scale, xd, up, - M, N, K, R, - a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), xd.stride(0), up.stride(0), up.stride(1), - QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), - BLOCK_R = block_r, R_TILES = triton.cdiv(R, block_r), - ) - return c - -@scaled_lora_mm_8bit.register_fake -def _(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: - return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype) - - -@triton.autotune(configs=_AUTOTUNE_CONFIGS, key=_LORA_AUTOTUNE_KEY, cache_results=True) -@triton.jit -def _rowcol_scaled_lora_mm_kernel( - a_ptr, b_ptr, c_ptr, scale_ptr, scale_n_ptr, xd_ptr, up_ptr, - M, N, K, R, - stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_xdm, stride_upr, stride_upn, - BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, - QUANTIZED_M, FLOAT: tl.constexpr, EVEN_K: tl.constexpr, MXFP8_MMA: tl.constexpr, - BLOCK_R: tl.constexpr, R_TILES: tl.constexpr, -): - accumulator, offs_cm, offs_cn = _mm_accumulate( - a_ptr, b_ptr, M, N, K, - stride_am, stride_ak, stride_bk, stride_bn, - BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K, GROUP_SIZE_M, - FLOAT, EVEN_K, MXFP8_MMA, - ) - - scale = tl.load(scale_ptr + offs_cm, mask=offs_cm < M, other=0.0) - scale_n = tl.load(scale_n_ptr + offs_cn, mask=offs_cn < N, other=0.0) - result = accumulator.to(tl.float32) * scale[:, None] * scale_n[None, :] - - result = _add_lora(result, xd_ptr, up_ptr, offs_cm, offs_cn, M, N, R, - stride_xdm, stride_upr, stride_upn, BLOCK_R, R_TILES) - result = result.to(c_ptr.dtype.element_ty) - - _store_c(c_ptr, result, offs_cm, offs_cn, M, N, stride_cm) - -announce_autotuning(_rowcol_scaled_lora_mm_kernel, name="8-bit row/column scaled LoRA matmul") +announce_autotuning(_mm_kernel, name="8-bit matmul") -@torch.library.custom_op("ot_quant::rowcol_scaled_lora_mm_8bit", mutates_args=()) -def rowcol_scaled_lora_mm_8bit(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: - R, block_r = _prepare_lora(a, b, xd, up) - b, c, scale, scale_n, M, N, K, FLOAT = _prepare_rowcol_scaled_mm(a, b, scale, scale_n, out_dtype) +#Opaque custom ops work around pytorch#164124: torch.compile otherwise absorbs the traced +#@triton.autotune kernels and freezes the config benchmarked for the first shape. Kept opaque, +#these bodies run eagerly, so Triton's autotuner selects per key and the JIT sees real sizes +@torch.library.custom_op("ot_quant::mm_8bit", mutates_args=()) +def mm_8bit(a: torch.Tensor, b: torch.Tensor, out_dtype: torch.dtype, + scale_m: torch.Tensor | None = None, scale_n: torch.Tensor | None = None, + lora_xd: torch.Tensor | None = None, lora_up: torch.Tensor | None = None) -> torch.Tensor: + #returns (a @ b) * scale_m * scale_n + lora_xd @ lora_up in out_dtype, each epilogue term optional + assert out_dtype.is_floating_point or (scale_m is None and scale_n is None and lora_xd is None), \ + "an integer out_dtype takes no epilogue" + a, b, c, M, N, K, FLOAT = _prepare_mm(a, b, out_dtype) + + if scale_m is not None: + scale_m = _prepare_scale(scale_m, M, "row of a") + if scale_n is not None: + scale_n = _prepare_scale(scale_n, N, "column of b") + + if lora_xd is not None: + R, block_r, r_tiles, stride_xdm, stride_upr, stride_upn = _prepare_lora(lora_xd, lora_up, M, N) + else: + assert lora_up is None, "lora_up must be passed together with lora_xd" + R = block_r = r_tiles = stride_xdm = stride_upr = stride_upn = None + #1D grid: the kernel derives pid_m/pid_n itself in grouped order for L2 reuse def grid(META): return (triton.cdiv(N, META['BLOCK_SIZE_N']) * triton.cdiv(M, META['BLOCK_SIZE_M']), ) - _rowcol_scaled_lora_mm_kernel[grid]( - a, b, c, scale, scale_n, xd, up, - M, N, K, R, - a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), xd.stride(0), up.stride(0), up.stride(1), - QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, EVEN_K = (K % 128 == 0), MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), - BLOCK_R = block_r, R_TILES = triton.cdiv(R, block_r), + _mm_kernel[grid]( + a, b, c, + M, N, K, + a.stride(0), a.stride(1), b.stride(0), b.stride(1), c.stride(0), + QUANTIZED_M = M.bit_length(), FLOAT = FLOAT, MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), + scale_m_ptr = scale_m, scale_n_ptr = scale_n, + xd_ptr = lora_xd, up_ptr = lora_up, R = R, stride_xdm = stride_xdm, stride_upr = stride_upr, stride_upn = stride_upn, + BLOCK_R = block_r, R_TILES = r_tiles, ) return c -@rowcol_scaled_lora_mm_8bit.register_fake -def _(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, scale_n: torch.Tensor, xd: torch.Tensor, up: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: +@mm_8bit.register_fake +def _(a: torch.Tensor, b: torch.Tensor, out_dtype: torch.dtype, + scale_m: torch.Tensor | None = None, scale_n: torch.Tensor | None = None, + lora_xd: torch.Tensor | None = None, lora_up: torch.Tensor | None = None) -> torch.Tensor: return a.new_empty((a.shape[0], b.shape[1]), dtype=out_dtype)