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/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/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/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..27f89c949 100644 --- a/modules/module/quantized/LinearGGUFA8.py +++ b/modules/module/quantized/LinearGGUFA8.py @@ -1,5 +1,7 @@ +from modules.module.quantized.mixin.LoRAFusableLinearMixin import LoRAFusableLinearMixin from modules.util.mm_8bit import mm_8bit as mm_8bit from modules.util.quantization_util import ( + quantize_axiswise, quantize_fp8_axiswise, quantize_int8_axiswise, ) @@ -14,70 +16,125 @@ 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 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) + #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 = 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 @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 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 = 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 @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 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 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)) -class LinearGGUFIntA8RequantFunction(torch.autograd.Function): + +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 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 +144,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..36180d5e6 100644 --- a/modules/module/quantized/LinearW8A8.py +++ b/modules/module/quantized/LinearW8A8.py @@ -1,10 +1,11 @@ - 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.quantization_util import ( dequantize, + quantize_axiswise, quantize_fp8_axiswise, quantize_fp8_tensorwise, quantize_int8_axiswise, @@ -16,40 +17,58 @@ @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 = 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 int8_backward_axiswise(output: Tensor, weight: Tensor, weight_scale: Tensor) -> Tensor: - output_8, output_scale = quantize_int8_axiswise(output, dim=-1) +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) #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) + return mm_8bit(output_8.contiguous(), weight, out_dtype=output.dtype, scale_m=weight_scale * output_scale) + + +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 = 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 @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_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 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 +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 +78,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 +86,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 +127,7 @@ class LinearW8A8( QuantizedModuleMixin, QuantizedLinearMixin, CompressedWeightMixin, + LoRAFusableLinearMixin, ): def __init__(self, dtype: torch.dtype, *args, **kwargs): super().__init__(*args, **kwargs) @@ -145,6 +192,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 +229,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 +238,20 @@ 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, 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") - 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, out_dtype=torch.int32), "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: 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 +260,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, 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") - 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, 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) + 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/trainer/BaseTrainer.py b/modules/trainer/BaseTrainer.py index 20fc5eb7a..87a96d7d0 100644 --- a/modules/trainer/BaseTrainer.py +++ b/modules/trainer/BaseTrainer.py @@ -97,7 +97,11 @@ def _start_tensorboard(self): if self.config.tensorboard_expose: tensorboard_args.append("--bind_all") - self.tensorboard_subprocess = subprocess.Popen(tensorboard_args) + # 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, + ) def _stop_tensorboard(self): self.tensorboard_subprocess.kill() diff --git a/modules/trainer/GenericTrainer.py b/modules/trainer/GenericTrainer.py index a6547afbe..61579f314 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 - # OT_DEBUG_PROFILES=1 dumps a CUDA memory snapshot for the first two steps, where the allocator is still # growing, and a profiler trace at steps 10 and 40, past compilation and warmup. _DEBUG_PROFILES = os.environ.get("OT_DEBUG_PROFILES") == "1" diff --git a/modules/ui/TrainUIController.py b/modules/ui/TrainUIController.py index 7aa550cee..895be42fa 100644 --- a/modules/ui/TrainUIController.py +++ b/modules/ui/TrainUIController.py @@ -108,8 +108,11 @@ def _start_always_on_tensorboard(self): if self.train_config.tensorboard_expose: tensorboard_args.append("--bind_all") + # discard the child's banner and notices; the UI already shows 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/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/mm_8bit.py b/modules/util/mm_8bit.py index a4e74f4bf..f13452191 100644 --- a/modules/util/mm_8bit.py +++ b/modules/util/mm_8bit.py @@ -1,15 +1,29 @@ +import torch + try: from modules.util.triton_mm_8bit import 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: + 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) - return torch._scaled_mm(a, b.T.contiguous().T, scale_a=one, scale_b=one) + res = torch._scaled_mm(a, b.T.contiguous().T, scale_a=one, scale_b=one, out_dtype=torch.float32) + + 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/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..4ece394ba 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 @@ -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/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/modules/util/triton_mm_8bit.py b/modules/util/triton_mm_8bit.py index 754d9ceda..97016205e 100644 --- a/modules/util/triton_mm_8bit.py +++ b/modules/util/triton_mm_8bit.py @@ -1,15 +1,13 @@ -#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 +#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 import torch @@ -25,59 +23,145 @@ def _prefer_mxfp8(device: torch.device) -> bool: return torch.version.cuda is not None and torch.cuda.get_device_capability(device)[0] >= 12 -@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"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.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 + +#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 = [ + #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) + '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 +#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 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, - MXFP8_MMA: 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, MXFP8_MMA: 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 @@ -94,11 +178,11 @@ def _mm_kernel( 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) - 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) + #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) if MXFP8_MMA: accumulator = tl.dot_scaled(a, a_scale, "e4m3", b, b_scale, "e4m3", acc=accumulator) @@ -110,11 +194,25 @@ def _mm_kernel( 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: +#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" assert a.dtype == b.dtype, "Incompatible dtypes" @@ -124,18 +222,142 @@ 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" + + #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 a, b, c, M, N, K, FLOAT + + +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] == 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 +#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(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( + 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=_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, 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, +): + 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, MXFP8_MMA, + ) + + 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) + + 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] + + 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, :] + + 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) + + _store_c(c_ptr, result.to(c_ptr.dtype.element_ty), 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, 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']), ) + 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, - MXFP8_MMA = FLOAT and _prefer_mxfp8(a.device), + 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 + +@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) diff --git a/modules/util/ui/pyside6_util.py b/modules/util/ui/pyside6_util.py index 07d77dce5..eb0f84853 100644 --- a/modules/util/ui/pyside6_util.py +++ b/modules/util/ui/pyside6_util.py @@ -1,4 +1,5 @@ import locale +import os import signal import sys from abc import ABCMeta @@ -20,6 +21,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) # QApplication initializes the C locale from the environment (setlocale(LC_ALL, "")), which sets LC_NUMERIC # to a locale whose decimal separator may be a comma. C libraries then misparse '.' floats: protobuf's upb diff --git a/scripts/util/import_util.py b/scripts/util/import_util.py index 150df4d8a..e269a0505 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,43 @@ 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. 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("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.*") + + # 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