diff --git a/setup.py b/setup.py index 4833cdfee..7a9f0f5bd 100644 --- a/setup.py +++ b/setup.py @@ -336,7 +336,7 @@ def _collect_package_files(*directories: str): "vector-quantize-pytorch>=1.27.15", "cryptography>=41.0.0", "torchcodec>=0.10.0", - "sdnq>=0.2.2", + "sdnq>=0.2.3", "aiosqlite>=0.19.0", "httpx>=0.28.0", "psutil>=5.9.0", diff --git a/simpletuner/helpers/models/krea2/quantized_loading.py b/simpletuner/helpers/models/krea2/quantized_loading.py index e731bb289..1a69c1c3b 100644 --- a/simpletuner/helpers/models/krea2/quantized_loading.py +++ b/simpletuner/helpers/models/krea2/quantized_loading.py @@ -129,9 +129,6 @@ def _materialize_krea2_meta_buffers(model: nn.Module) -> None: def _load_sdnq_training_symbols(): try: - from simpletuner.helpers.training.sdnq_compat import apply_sdnq_checkpointed_backward_fix - - apply_sdnq_checkpointed_backward_fix() from sdnq.dequantizer import SDNQDequantizer from sdnq.layers import get_sdnq_wrapper_class from sdnq.training.forward import get_forward_func diff --git a/simpletuner/helpers/models/z_image/quantized_loading.py b/simpletuner/helpers/models/z_image/quantized_loading.py index d82344cd3..81ac3ba2d 100644 --- a/simpletuner/helpers/models/z_image/quantized_loading.py +++ b/simpletuner/helpers/models/z_image/quantized_loading.py @@ -118,9 +118,6 @@ def _materialize_zimage_meta_buffers(model: nn.Module) -> None: def _load_sdnq_training_symbols(): try: - from simpletuner.helpers.training.sdnq_compat import apply_sdnq_checkpointed_backward_fix - - apply_sdnq_checkpointed_backward_fix() from sdnq.dequantizer import SDNQDequantizer from sdnq.layers import get_sdnq_wrapper_class from sdnq.training.forward import get_forward_func diff --git a/simpletuner/helpers/training/quantisation/__init__.py b/simpletuner/helpers/training/quantisation/__init__.py index 9ff896ede..72a13d0a4 100644 --- a/simpletuner/helpers/training/quantisation/__init__.py +++ b/simpletuner/helpers/training/quantisation/__init__.py @@ -1058,10 +1058,6 @@ def _sdnq_model( # Silence sdnq startup logs logging.getLogger("sdnq").setLevel(logging.WARNING) import sdnq.common as sdnq_common - - from simpletuner.helpers.training.sdnq_compat import apply_sdnq_checkpointed_backward_fix - - apply_sdnq_checkpointed_backward_fix(logger) from sdnq.training import sdnq_training_post_load_quant except ImportError as e: raise ImportError(f"To use SDNQ, please install the sdnq library: `pip install sdnq`: {e}") diff --git a/simpletuner/helpers/training/sdnq_compat.py b/simpletuner/helpers/training/sdnq_compat.py deleted file mode 100644 index 54f10b7b4..000000000 --- a/simpletuner/helpers/training/sdnq_compat.py +++ /dev/null @@ -1,921 +0,0 @@ -from __future__ import annotations - -import importlib -import logging -from importlib import metadata -from inspect import signature -from typing import Any - -import torch - -_PATCH_MARKER = "_simpletuner_fd6d7e0_checkpoint_patch" - - -def _sdnq_version() -> str | None: - try: - return metadata.version("sdnq") - except metadata.PackageNotFoundError: - return None - - -def _version_tuple(version: str) -> tuple[int, ...]: - parts: list[int] = [] - for part in version.split("."): - number = "" - for char in part: - if not char.isdigit(): - break - number += char - if not number: - break - parts.append(int(number)) - return tuple(parts) - - -def _has_upstream_checkpoint_fix() -> bool: - module = importlib.import_module("sdnq.training.layers.linear.linear_int8.linear_int8_ckpt") - return "input_shape" in signature(module.int8_matmul_backward_ckpt).parameters - - -def _output_shape(grad_output: torch.Tensor, input: torch.Tensor | None, input_shape: torch.Size | None) -> list[int]: - output_shape = list(grad_output.shape) - output_shape[-1] = input_shape[-1] if input_shape is not None else input.shape[-1] - return output_shape - - -def _optional_tensor(value: torch.Tensor | None, placeholder: torch.Tensor) -> torch.Tensor: - return value if value is not None else placeholder - - -def _restore_optional(value: torch.Tensor, present: bool) -> torch.Tensor | None: - return value if present else None - - -def _patch_int8_static(module: Any) -> None: - def int8_matmul_backward_ckpt( - grad_output, - input, - weight, - input_scale, - scale, - bias=None, - svd_up=None, - svd_down=None, - zero_point=None, - hadamard=None, - input_shape=None, - do_grad_input=True, - do_grad_weight=True, - do_grad_bias=True, - ): - grad_input = grad_weight = grad_bias = None - output_shape = _output_shape(grad_output, input, input_shape) - grad_output = grad_output.flatten(0, -2) - if do_grad_input: - dequantized_weight = ( - module.dequantize_symmetric_compiled(weight, scale) - if zero_point is None - else module.dequantize_asymmetric_compiled(weight, scale, zero_point) - ) - grad_input = module.int8_matmul_dynamic( - grad_output, - dequantized_weight, - svd_up=svd_up, - svd_down=svd_down, - hadamard=hadamard, - output_shape=output_shape, - do_input_reshape=False, - ) - if do_grad_weight: - grad_weight = module.int8_matmul( - grad_output.t(), - input, - input_scale, - hadamard=hadamard, - output_shape=None, - do_input_reshape=False, - do_transpose=False, - ) - if do_grad_bias and bias is not None: - grad_bias = grad_output.sum(dim=0) - return grad_input, grad_weight, grad_bias - - class INT8MatmulBackwardCKPT(torch.autograd.Function): - @staticmethod - def forward(ctx, input, weight, bias=None): - if weight.sdnq_dequantizer.use_hadamard: - hadamard = module.get_hadamard( - weight.sdnq_dequantizer.hadamard_group_size, - dtype=input.dtype, - device=input.device, - ) - else: - hadamard = None - - result = module.int8_matmul( - input, - weight.weight, - weight.scale, - bias=bias, - svd_up=weight.svd_up, - svd_down=weight.svd_down, - zero_point=weight.zero_point, - hadamard=hadamard, - do_transpose=True, - ) - if ctx.needs_input_grad[1]: - new_input, input_scale = module.get_int8_matmul_backward_inputs(input, hadamard) - else: - new_input = input_scale = None - placeholder = input.new_empty(0) - ctx.has_weight_grad_inputs = new_input is not None - ctx.has_bias = bias is not None - ctx.save_for_backward( - _optional_tensor(new_input, placeholder), - weight, - _optional_tensor(input_scale, placeholder), - _optional_tensor(bias, placeholder), - ) - ctx.input_shape = input.shape - return result - - @staticmethod - def backward(ctx, grad_output): - input, weight, input_scale, bias = ctx.saved_tensors - input = _restore_optional(input, ctx.has_weight_grad_inputs) - input_scale = _restore_optional(input_scale, ctx.has_weight_grad_inputs) - bias = _restore_optional(bias, ctx.has_bias) - if weight.sdnq_dequantizer.use_hadamard: - hadamard = module.get_hadamard( - weight.sdnq_dequantizer.hadamard_group_size, - dtype=grad_output.dtype, - device=grad_output.device, - ) - else: - hadamard = None - return module.int8_matmul_backward_ckpt( - grad_output, - input, - weight.weight, - input_scale, - weight.scale, - bias=bias, - svd_up=weight.svd_up, - svd_down=weight.svd_down, - zero_point=weight.zero_point, - hadamard=hadamard, - input_shape=ctx.input_shape, - do_grad_input=ctx.needs_input_grad[0], - do_grad_weight=ctx.needs_input_grad[1], - do_grad_bias=ctx.needs_input_grad[2], - ) - - module.int8_matmul_backward_ckpt = int8_matmul_backward_ckpt - module.INT8MatmulBackwardCKPT = INT8MatmulBackwardCKPT - module.int8_matmul_with_backward_ckpt = INT8MatmulBackwardCKPT.apply - module.int8_matmul_backward_ckpt.__dict__[_PATCH_MARKER] = True - - -def _patch_int8_dynamic(module: Any) -> None: - def get_int8_matmul_dynamic_backward_inputs(input, weight, hadamard, do_grad_weight=True): - weight, scale = module.quantize_int_mm(weight.to(dtype=torch.float32), dim=0) - if do_grad_weight: - input, input_scale = module.quantize_int_mm( - input.flatten(0, -2).to(dtype=torch.float32), - dim=0, - hadamard=hadamard, - ) - return input, weight, input_scale, scale - return None, weight, None, scale - - def int8_matmul_dynamic_backward_ckpt( - grad_output, - input, - weight, - input_scale, - weight_scale, - bias=None, - svd_up=None, - svd_down=None, - hadamard=None, - input_shape=None, - do_grad_input=True, - do_grad_weight=True, - do_grad_bias=True, - ): - grad_input = grad_weight = grad_bias = None - output_shape = _output_shape(grad_output, input, input_shape) - grad_output = grad_output.flatten(0, -2) - if do_grad_input: - grad_input = module.int8_matmul( - grad_output, - weight, - weight_scale, - svd_up=svd_up, - svd_down=svd_down, - hadamard=hadamard, - output_shape=output_shape, - do_input_reshape=False, - do_transpose=False, - ) - if do_grad_weight: - grad_weight = module.int8_matmul( - grad_output.t(), - input, - input_scale, - hadamard=hadamard, - output_shape=None, - do_input_reshape=False, - do_transpose=False, - ) - if do_grad_bias and bias is not None: - grad_bias = grad_output.sum(dim=0) - return grad_input, grad_weight, grad_bias - - class INT8MatmulDynamicBackwardCKPT(torch.autograd.Function): - @staticmethod - def forward(ctx, input, weight, bias=None): - if isinstance(weight, module.SDNQTensor): - svd_up, svd_down = weight.svd_up, weight.svd_down - ctx.use_hadamard = weight.sdnq_dequantizer.use_hadamard - ctx.hadamard_group_size = weight.sdnq_dequantizer.hadamard_group_size - weight = weight.dequantize(non_svd=True, non_hadamard=True) - else: - svd_up, svd_down = None, None - ctx.use_hadamard = False - ctx.hadamard_group_size = 256 - hadamard = ( - module.get_hadamard(ctx.hadamard_group_size, dtype=input.dtype, device=input.device) - if ctx.use_hadamard - else None - ) - result = module.int8_matmul_dynamic( - input, - weight, - bias=bias, - svd_up=svd_up, - svd_down=svd_down, - hadamard=hadamard, - ) - new_input, new_weight, input_scale, weight_scale = module.get_int8_matmul_dynamic_backward_inputs( - input, - weight, - hadamard, - do_grad_weight=ctx.needs_input_grad[1], - ) - placeholder = input.new_empty(0) - ctx.has_weight_grad_inputs = new_input is not None - ctx.has_bias = bias is not None - ctx.has_svd_up = svd_up is not None - ctx.has_svd_down = svd_down is not None - ctx.save_for_backward( - _optional_tensor(new_input, placeholder), - new_weight, - _optional_tensor(input_scale, placeholder), - weight_scale, - _optional_tensor(bias, placeholder), - _optional_tensor(svd_up, placeholder), - _optional_tensor(svd_down, placeholder), - ) - ctx.input_shape = input.shape - return result - - @staticmethod - def backward(ctx, grad_output): - input, weight, input_scale, weight_scale, bias, svd_up, svd_down = ctx.saved_tensors - input = _restore_optional(input, ctx.has_weight_grad_inputs) - input_scale = _restore_optional(input_scale, ctx.has_weight_grad_inputs) - bias = _restore_optional(bias, ctx.has_bias) - svd_up = _restore_optional(svd_up, ctx.has_svd_up) - svd_down = _restore_optional(svd_down, ctx.has_svd_down) - hadamard = ( - module.get_hadamard(ctx.hadamard_group_size, dtype=grad_output.dtype, device=grad_output.device) - if ctx.use_hadamard - else None - ) - return module.int8_matmul_dynamic_backward_ckpt( - grad_output, - input, - weight, - input_scale, - weight_scale, - bias=bias, - svd_up=svd_up, - svd_down=svd_down, - hadamard=hadamard, - input_shape=ctx.input_shape, - do_grad_input=ctx.needs_input_grad[0], - do_grad_weight=ctx.needs_input_grad[1], - do_grad_bias=ctx.needs_input_grad[2], - ) - - module.get_int8_matmul_dynamic_backward_inputs = module.compile_func(get_int8_matmul_dynamic_backward_inputs) - module.int8_matmul_dynamic_backward_ckpt = int8_matmul_dynamic_backward_ckpt - module.INT8MatmulDynamicBackwardCKPT = INT8MatmulDynamicBackwardCKPT - module.int8_matmul_dynamic_with_backward_ckpt = INT8MatmulDynamicBackwardCKPT.apply - module.int8_matmul_dynamic_backward_ckpt.__dict__[_PATCH_MARKER] = True - - -def _patch_fp_static(module: Any, *, dtype_name: str, matmul_dtype: str) -> None: - matmul = getattr(module, f"{dtype_name}_matmul") - matmul_dynamic = getattr(module, f"{dtype_name}_matmul_dynamic") - - def backward_ckpt( - grad_output, - input, - weight, - input_scale, - scale, - bias=None, - svd_up=None, - svd_down=None, - hadamard=None, - input_shape=None, - do_grad_input=True, - do_grad_weight=True, - do_grad_bias=True, - ): - grad_input = grad_weight = grad_bias = None - output_shape = _output_shape(grad_output, input, input_shape) - grad_output = grad_output.flatten(0, -2) - if do_grad_input: - grad_input = matmul_dynamic( - grad_output, - module.dequantize_symmetric_compiled(weight, scale), - svd_up=svd_up, - svd_down=svd_down, - hadamard=hadamard, - output_shape=output_shape, - do_input_reshape=False, - ) - if do_grad_weight: - grad_weight = matmul( - grad_output.t(), - input, - input_scale, - hadamard=hadamard, - output_shape=None, - do_input_reshape=False, - do_transpose=False, - ) - if do_grad_bias and bias is not None: - grad_bias = grad_output.sum(dim=0) - return grad_input, grad_weight, grad_bias - - class FPMatmulBackwardCKPT(torch.autograd.Function): - @staticmethod - def forward(ctx, input, weight, bias=None): - hadamard = ( - module.get_hadamard( - weight.sdnq_dequantizer.hadamard_group_size, - dtype=input.dtype, - device=input.device, - ) - if weight.sdnq_dequantizer.use_hadamard - else None - ) - result = matmul( - input, - weight.weight, - weight.scale, - bias=bias, - svd_up=weight.svd_up, - svd_down=weight.svd_down, - hadamard=hadamard, - do_transpose=True, - ) - if ctx.needs_input_grad[1]: - new_input, input_scale = module.quantize_fp_mm( - input.flatten(0, -2).to(dtype=torch.float32), - dim=0, - hadamard=hadamard, - matmul_dtype=matmul_dtype, - ) - else: - new_input = input_scale = None - placeholder = input.new_empty(0) - ctx.has_weight_grad_inputs = new_input is not None - ctx.has_bias = bias is not None - ctx.save_for_backward( - _optional_tensor(new_input, placeholder), - weight, - _optional_tensor(input_scale, placeholder), - _optional_tensor(bias, placeholder), - ) - ctx.input_shape = input.shape - return result - - @staticmethod - def backward(ctx, grad_output): - input, weight, input_scale, bias = ctx.saved_tensors - input = _restore_optional(input, ctx.has_weight_grad_inputs) - input_scale = _restore_optional(input_scale, ctx.has_weight_grad_inputs) - bias = _restore_optional(bias, ctx.has_bias) - hadamard = ( - module.get_hadamard( - weight.sdnq_dequantizer.hadamard_group_size, - dtype=grad_output.dtype, - device=grad_output.device, - ) - if weight.sdnq_dequantizer.use_hadamard - else None - ) - return getattr(module, f"{dtype_name}_matmul_backward_ckpt")( - grad_output, - input, - weight.weight, - input_scale, - weight.scale, - bias=bias, - svd_up=weight.svd_up, - svd_down=weight.svd_down, - hadamard=hadamard, - input_shape=ctx.input_shape, - do_grad_input=ctx.needs_input_grad[0], - do_grad_weight=ctx.needs_input_grad[1], - do_grad_bias=ctx.needs_input_grad[2], - ) - - setattr(module, f"{dtype_name}_matmul_backward_ckpt", backward_ckpt) - setattr(module, f"{dtype_name.upper()}MatmulBackwardCKPT", FPMatmulBackwardCKPT) - setattr(module, f"{dtype_name}_matmul_with_backward_ckpt", FPMatmulBackwardCKPT.apply) - getattr(module, f"{dtype_name}_matmul_backward_ckpt").__dict__[_PATCH_MARKER] = True - - -def _patch_fp_dynamic(module: Any, *, dtype_name: str, matmul_dtype: str) -> None: - matmul = getattr(module, f"{dtype_name}_matmul") - matmul_dynamic = getattr(module, f"{dtype_name}_matmul_dynamic") - - def dynamic_backward_inputs(input, weight, hadamard, do_grad_weight=True): - new_weight, weight_scale = module.quantize_fp_mm( - weight.to(dtype=torch.float32), - dim=0, - matmul_dtype=matmul_dtype, - ) - if do_grad_weight: - new_input, input_scale = module.quantize_fp_mm( - input.flatten(0, -2).to(dtype=torch.float32), - dim=0, - hadamard=hadamard, - matmul_dtype=matmul_dtype, - ) - return new_input, new_weight, input_scale, weight_scale - return None, new_weight, None, weight_scale - - def dynamic_backward_ckpt( - grad_output, - input, - weight, - input_scale, - weight_scale, - bias=None, - svd_up=None, - svd_down=None, - hadamard=None, - input_shape=None, - do_grad_input=True, - do_grad_weight=True, - do_grad_bias=True, - ): - grad_input = grad_weight = grad_bias = None - output_shape = _output_shape(grad_output, input, input_shape) - grad_output = grad_output.flatten(0, -2) - if do_grad_input: - grad_input = matmul( - grad_output, - weight, - weight_scale, - svd_up=svd_up, - svd_down=svd_down, - hadamard=hadamard, - output_shape=output_shape, - do_input_reshape=False, - do_transpose=False, - ) - if do_grad_weight: - grad_weight = matmul( - grad_output.t(), - input, - input_scale, - hadamard=hadamard, - output_shape=None, - do_input_reshape=False, - do_transpose=False, - ) - if do_grad_bias and bias is not None: - grad_bias = grad_output.sum(dim=0) - return grad_input, grad_weight, grad_bias - - class FPMatmulDynamicBackwardCKPT(torch.autograd.Function): - @staticmethod - def forward(ctx, input, weight, bias=None): - if isinstance(weight, module.SDNQTensor): - svd_up, svd_down = weight.svd_up, weight.svd_down - ctx.use_hadamard = weight.sdnq_dequantizer.use_hadamard - ctx.hadamard_group_size = weight.sdnq_dequantizer.hadamard_group_size - weight = weight.dequantize(non_svd=True, non_hadamard=True) - else: - svd_up, svd_down = None, None - ctx.use_hadamard = False - ctx.hadamard_group_size = 256 - hadamard = ( - module.get_hadamard(ctx.hadamard_group_size, dtype=input.dtype, device=input.device) - if ctx.use_hadamard - else None - ) - result = matmul_dynamic( - input, - weight, - bias=bias, - svd_up=svd_up, - svd_down=svd_down, - hadamard=hadamard, - ) - new_input, new_weight, input_scale, weight_scale = getattr( - module, - f"get_{dtype_name}_matmul_dynamic_backward_inputs", - )( - input, - weight, - hadamard, - do_grad_weight=ctx.needs_input_grad[1], - ) - placeholder = input.new_empty(0) - ctx.has_weight_grad_inputs = new_input is not None - ctx.has_bias = bias is not None - ctx.has_svd_up = svd_up is not None - ctx.has_svd_down = svd_down is not None - ctx.save_for_backward( - _optional_tensor(new_input, placeholder), - new_weight, - _optional_tensor(input_scale, placeholder), - weight_scale, - _optional_tensor(bias, placeholder), - _optional_tensor(svd_up, placeholder), - _optional_tensor(svd_down, placeholder), - ) - ctx.input_shape = input.shape - return result - - @staticmethod - def backward(ctx, grad_output): - input, weight, input_scale, weight_scale, bias, svd_up, svd_down = ctx.saved_tensors - input = _restore_optional(input, ctx.has_weight_grad_inputs) - input_scale = _restore_optional(input_scale, ctx.has_weight_grad_inputs) - bias = _restore_optional(bias, ctx.has_bias) - svd_up = _restore_optional(svd_up, ctx.has_svd_up) - svd_down = _restore_optional(svd_down, ctx.has_svd_down) - hadamard = ( - module.get_hadamard(ctx.hadamard_group_size, dtype=grad_output.dtype, device=grad_output.device) - if ctx.use_hadamard - else None - ) - return getattr(module, f"{dtype_name}_matmul_dynamic_backward_ckpt")( - grad_output, - input, - weight, - input_scale, - weight_scale, - bias=bias, - svd_up=svd_up, - svd_down=svd_down, - hadamard=hadamard, - input_shape=ctx.input_shape, - do_grad_input=ctx.needs_input_grad[0], - do_grad_weight=ctx.needs_input_grad[1], - do_grad_bias=ctx.needs_input_grad[2], - ) - - setattr(module, f"get_{dtype_name}_matmul_dynamic_backward_inputs", module.compile_func(dynamic_backward_inputs)) - setattr(module, f"{dtype_name}_matmul_dynamic_backward_ckpt", dynamic_backward_ckpt) - setattr(module, f"{dtype_name.upper()}MatmulDynamicBackwardCKPT", FPMatmulDynamicBackwardCKPT) - setattr(module, f"{dtype_name}_matmul_dynamic_with_backward_ckpt", FPMatmulDynamicBackwardCKPT.apply) - getattr(module, f"{dtype_name}_matmul_dynamic_backward_ckpt").__dict__[_PATCH_MARKER] = True - - -def _patch_uint8_static(module: Any) -> None: - def uint8_matmul_backward_ckpt( - grad_output, - input, - weight, - input_scale, - scale, - input_zero_point, - zero_point, - bias=None, - svd_up=None, - svd_down=None, - hadamard=None, - input_shape=None, - do_grad_input=True, - do_grad_weight=True, - do_grad_bias=True, - ): - grad_input = grad_weight = grad_bias = None - output_shape = _output_shape(grad_output, input, input_shape) - grad_output = grad_output.flatten(0, -2) - if do_grad_input: - dequantized_weight = ( - module.dequantize_symmetric_compiled(weight, scale) - if zero_point is None - else module.dequantize_asymmetric_compiled(weight, scale, zero_point) - ) - grad_input = module.uint8_matmul_dynamic( - grad_output, - dequantized_weight, - svd_up=svd_up, - svd_down=svd_down, - hadamard=hadamard, - output_shape=output_shape, - do_input_reshape=False, - ) - if do_grad_weight: - grad_weight = module.uint8_matmul( - grad_output.t(), - input, - input_scale, - input_zero_point, - hadamard=hadamard, - output_shape=None, - do_input_reshape=False, - do_transpose=False, - ) - if do_grad_bias and bias is not None: - grad_bias = grad_output.sum(dim=0) - return grad_input, grad_weight, grad_bias - - class UINT8MatmulBackwardCKPT(torch.autograd.Function): - @staticmethod - def forward(ctx, input, weight, bias=None): - hadamard = ( - module.get_hadamard( - weight.sdnq_dequantizer.hadamard_group_size, - dtype=input.dtype, - device=input.device, - ) - if weight.sdnq_dequantizer.use_hadamard - else None - ) - result = module.uint8_matmul( - input, - weight.weight, - weight.scale, - weight.zero_point, - bias=bias, - svd_up=weight.svd_up, - svd_down=weight.svd_down, - hadamard=hadamard, - do_transpose=True, - ) - if ctx.needs_input_grad[1]: - new_input, input_scale, input_zero_point = module.get_uint8_matmul_backward_inputs(input, hadamard) - else: - new_input = input_scale = input_zero_point = None - placeholder = input.new_empty(0) - ctx.has_weight_grad_inputs = new_input is not None - ctx.has_bias = bias is not None - ctx.save_for_backward( - _optional_tensor(new_input, placeholder), - weight, - _optional_tensor(input_scale, placeholder), - _optional_tensor(input_zero_point, placeholder), - _optional_tensor(bias, placeholder), - ) - ctx.input_shape = input.shape - return result - - @staticmethod - def backward(ctx, grad_output): - input, weight, input_scale, input_zero_point, bias = ctx.saved_tensors - input = _restore_optional(input, ctx.has_weight_grad_inputs) - input_scale = _restore_optional(input_scale, ctx.has_weight_grad_inputs) - input_zero_point = _restore_optional(input_zero_point, ctx.has_weight_grad_inputs) - bias = _restore_optional(bias, ctx.has_bias) - hadamard = ( - module.get_hadamard( - weight.sdnq_dequantizer.hadamard_group_size, - dtype=grad_output.dtype, - device=grad_output.device, - ) - if weight.sdnq_dequantizer.use_hadamard - else None - ) - return module.uint8_matmul_backward_ckpt( - grad_output, - input, - weight.weight, - input_scale, - weight.scale, - input_zero_point, - weight.zero_point, - bias=bias, - svd_up=weight.svd_up, - svd_down=weight.svd_down, - hadamard=hadamard, - input_shape=ctx.input_shape, - do_grad_input=ctx.needs_input_grad[0], - do_grad_weight=ctx.needs_input_grad[1], - do_grad_bias=ctx.needs_input_grad[2], - ) - - module.uint8_matmul_backward_ckpt = uint8_matmul_backward_ckpt - module.UINT8MatmulBackwardCKPT = UINT8MatmulBackwardCKPT - module.uint8_matmul_with_backward_ckpt = UINT8MatmulBackwardCKPT.apply - module.uint8_matmul_backward_ckpt.__dict__[_PATCH_MARKER] = True - - -def _patch_uint8_dynamic(module: Any) -> None: - def get_uint8_matmul_dynamic_backward_inputs(input, weight, hadamard, do_grad_weight=True): - weight, scale, zero_point = module.quantize_uint_mm(weight.to(dtype=torch.float32), dim=0) - if do_grad_weight: - input, input_scale, input_zero_point = module.quantize_uint_mm( - input.flatten(0, -2).to(dtype=torch.float32), - dim=0, - hadamard=hadamard, - ) - return input, weight, input_scale, scale, input_zero_point, zero_point - return None, weight, None, scale, None, zero_point - - def uint8_matmul_dynamic_backward_ckpt( - grad_output, - input, - weight, - input_scale, - weight_scale, - input_zero_point, - weight_zero_point, - bias=None, - svd_up=None, - svd_down=None, - hadamard=None, - input_shape=None, - do_grad_input=True, - do_grad_weight=True, - do_grad_bias=True, - ): - grad_input = grad_weight = grad_bias = None - output_shape = _output_shape(grad_output, input, input_shape) - grad_output = grad_output.flatten(0, -2) - if do_grad_input: - grad_input = module.uint8_matmul( - grad_output, - weight, - weight_scale, - weight_zero_point, - svd_up=svd_up, - svd_down=svd_down, - hadamard=hadamard, - output_shape=output_shape, - do_input_reshape=False, - do_transpose=False, - ) - if do_grad_weight: - grad_weight = module.uint8_matmul( - grad_output.t(), - input, - input_scale, - input_zero_point, - hadamard=hadamard, - output_shape=None, - do_input_reshape=False, - do_transpose=False, - ) - if do_grad_bias and bias is not None: - grad_bias = grad_output.sum(dim=0) - return grad_input, grad_weight, grad_bias - - class UINT8MatmulDynamicBackwardCKPT(torch.autograd.Function): - @staticmethod - def forward(ctx, input, weight, bias=None): - if isinstance(weight, module.SDNQTensor): - svd_up, svd_down = weight.svd_up, weight.svd_down - ctx.use_hadamard = weight.sdnq_dequantizer.use_hadamard - ctx.hadamard_group_size = weight.sdnq_dequantizer.hadamard_group_size - weight = weight.dequantize(non_svd=True, non_hadamard=True) - else: - svd_up, svd_down = None, None - ctx.use_hadamard = False - ctx.hadamard_group_size = 256 - hadamard = ( - module.get_hadamard(ctx.hadamard_group_size, dtype=input.dtype, device=input.device) - if ctx.use_hadamard - else None - ) - result = module.uint8_matmul_dynamic( - input, - weight, - bias=bias, - svd_up=svd_up, - svd_down=svd_down, - hadamard=hadamard, - ) - new_input, new_weight, input_scale, weight_scale, input_zero_point, weight_zero_point = ( - module.get_uint8_matmul_dynamic_backward_inputs( - input, - weight, - hadamard, - do_grad_weight=ctx.needs_input_grad[1], - ) - ) - placeholder = input.new_empty(0) - ctx.has_weight_grad_inputs = new_input is not None - ctx.has_bias = bias is not None - ctx.has_svd_up = svd_up is not None - ctx.has_svd_down = svd_down is not None - ctx.save_for_backward( - _optional_tensor(new_input, placeholder), - new_weight, - _optional_tensor(input_scale, placeholder), - weight_scale, - _optional_tensor(input_zero_point, placeholder), - weight_zero_point, - _optional_tensor(bias, placeholder), - _optional_tensor(svd_up, placeholder), - _optional_tensor(svd_down, placeholder), - ) - ctx.input_shape = input.shape - return result - - @staticmethod - def backward(ctx, grad_output): - input, weight, input_scale, weight_scale, input_zero_point, weight_zero_point, bias, svd_up, svd_down = ( - ctx.saved_tensors - ) - input = _restore_optional(input, ctx.has_weight_grad_inputs) - input_scale = _restore_optional(input_scale, ctx.has_weight_grad_inputs) - input_zero_point = _restore_optional(input_zero_point, ctx.has_weight_grad_inputs) - bias = _restore_optional(bias, ctx.has_bias) - svd_up = _restore_optional(svd_up, ctx.has_svd_up) - svd_down = _restore_optional(svd_down, ctx.has_svd_down) - hadamard = ( - module.get_hadamard(ctx.hadamard_group_size, dtype=grad_output.dtype, device=grad_output.device) - if ctx.use_hadamard - else None - ) - return module.uint8_matmul_dynamic_backward_ckpt( - grad_output, - input, - weight, - input_scale, - weight_scale, - input_zero_point, - weight_zero_point, - bias=bias, - svd_up=svd_up, - svd_down=svd_down, - hadamard=hadamard, - input_shape=ctx.input_shape, - do_grad_input=ctx.needs_input_grad[0], - do_grad_weight=ctx.needs_input_grad[1], - do_grad_bias=ctx.needs_input_grad[2], - ) - - module.get_uint8_matmul_dynamic_backward_inputs = module.compile_func(get_uint8_matmul_dynamic_backward_inputs) - module.uint8_matmul_dynamic_backward_ckpt = uint8_matmul_dynamic_backward_ckpt - module.UINT8MatmulDynamicBackwardCKPT = UINT8MatmulDynamicBackwardCKPT - module.uint8_matmul_dynamic_with_backward_ckpt = UINT8MatmulDynamicBackwardCKPT.apply - module.uint8_matmul_dynamic_backward_ckpt.__dict__[_PATCH_MARKER] = True - - -def apply_sdnq_checkpointed_backward_fix(logger: logging.Logger | None = None) -> bool: - """Apply Disty0/sdnq fd6d7e0 when the installed SDNQ wheel predates it.""" - - version = _sdnq_version() - if version is None or _version_tuple(version) < (0, 2, 2): - return False - - try: - int8_static = importlib.import_module("sdnq.training.layers.linear.linear_int8.linear_int8_ckpt") - if getattr(int8_static.int8_matmul_backward_ckpt, _PATCH_MARKER, False) or _has_upstream_checkpoint_fix(): - return False - - _patch_int8_static(int8_static) - _patch_int8_dynamic(importlib.import_module("sdnq.training.layers.linear.linear_int8.linear_int8_dynamic_ckpt")) - _patch_uint8_static(importlib.import_module("sdnq.training.layers.linear.linear_uint8.linear_uint8_ckpt")) - _patch_uint8_dynamic(importlib.import_module("sdnq.training.layers.linear.linear_uint8.linear_uint8_dynamic_ckpt")) - _patch_fp_static( - importlib.import_module("sdnq.training.layers.linear.linear_fp8.linear_fp8_ckpt"), - dtype_name="fp8", - matmul_dtype="float8_e4m3fn", - ) - _patch_fp_dynamic( - importlib.import_module("sdnq.training.layers.linear.linear_fp8.linear_fp8_dynamic_ckpt"), - dtype_name="fp8", - matmul_dtype="float8_e4m3fn", - ) - _patch_fp_static( - importlib.import_module("sdnq.training.layers.linear.linear_fp16.linear_fp16_ckpt"), - dtype_name="fp16", - matmul_dtype="float16", - ) - _patch_fp_dynamic( - importlib.import_module("sdnq.training.layers.linear.linear_fp16.linear_fp16_dynamic_ckpt"), - dtype_name="fp16", - matmul_dtype="float16", - ) - except ModuleNotFoundError: - return False - - if logger is not None: - logger.info("Applied SDNQ checkpointed backward compatibility patch for frozen quantized weights.") - return True diff --git a/tests/test_sdnq_compat.py b/tests/test_sdnq_compat.py deleted file mode 100644 index 61db86c80..000000000 --- a/tests/test_sdnq_compat.py +++ /dev/null @@ -1,222 +0,0 @@ -from __future__ import annotations - -import sys -import types -import unittest -from inspect import signature -from unittest.mock import patch - -import torch - -from simpletuner.helpers.training.sdnq_compat import apply_sdnq_checkpointed_backward_fix - -MODULE_NAMES = ( - "sdnq.training.layers.linear.linear_int8.linear_int8_ckpt", - "sdnq.training.layers.linear.linear_int8.linear_int8_dynamic_ckpt", - "sdnq.training.layers.linear.linear_uint8.linear_uint8_ckpt", - "sdnq.training.layers.linear.linear_uint8.linear_uint8_dynamic_ckpt", - "sdnq.training.layers.linear.linear_fp8.linear_fp8_ckpt", - "sdnq.training.layers.linear.linear_fp8.linear_fp8_dynamic_ckpt", - "sdnq.training.layers.linear.linear_fp16.linear_fp16_ckpt", - "sdnq.training.layers.linear.linear_fp16.linear_fp16_dynamic_ckpt", -) - - -def _install_module(name: str) -> types.ModuleType: - module = types.ModuleType(name) - sys.modules[name] = module - return module - - -def _matmul(input: torch.Tensor, weight: torch.Tensor, scale: torch.Tensor | None = None, **kwargs) -> torch.Tensor: - output = input.flatten(0, -2).to(dtype=torch.float32).matmul(weight.to(dtype=torch.float32).t()) - output_shape = kwargs.get("output_shape") - if output_shape is not None: - output = output.reshape(output_shape) - return output.to(dtype=input.dtype) - - -def _dynamic_matmul(input: torch.Tensor, weight: torch.Tensor, **kwargs) -> torch.Tensor: - output = input.flatten(0, -2).to(dtype=torch.float32).matmul(weight.to(dtype=torch.float32)) - output_shape = kwargs.get("output_shape") - if output_shape is not None: - output = output.reshape(output_shape) - return output.to(dtype=input.dtype) - - -def _quantize_int_mm(input: torch.Tensor, **kwargs): - return input.round().to(torch.int8), torch.ones((input.shape[0], 1), dtype=torch.float32, device=input.device) - - -def _quantize_uint_mm(input: torch.Tensor, **kwargs): - return ( - input.round().clamp_min(0).to(torch.uint8), - torch.ones((input.shape[0], 1), dtype=torch.float32, device=input.device), - torch.zeros((input.shape[0], 1), dtype=torch.float32, device=input.device), - ) - - -def _quantize_fp_mm(input: torch.Tensor, **kwargs): - return input.to(torch.float16), torch.ones((input.shape[0], 1), dtype=torch.float32, device=input.device) - - -def _install_fake_sdnq_modules(*, upstream_fixed: bool = False) -> None: - for package_name in ( - "sdnq", - "sdnq.training", - "sdnq.training.layers", - "sdnq.training.layers.linear", - "sdnq.training.layers.linear.linear_int8", - "sdnq.training.layers.linear.linear_uint8", - "sdnq.training.layers.linear.linear_fp8", - "sdnq.training.layers.linear.linear_fp16", - ): - _install_module(package_name) - - for module_name in MODULE_NAMES: - module = _install_module(module_name) - module.compile_func = lambda func: func - module.SDNQTensor = torch.Tensor - module.get_hadamard = lambda group_size, dtype, device: torch.eye(group_size, dtype=dtype, device=device) - module.dequantize_symmetric_compiled = lambda weight, scale: weight.to(dtype=torch.float32) - module.dequantize_asymmetric_compiled = lambda weight, scale, zero_point: weight.to(dtype=torch.float32) - - int8_static = sys.modules["sdnq.training.layers.linear.linear_int8.linear_int8_ckpt"] - int8_static.int8_matmul = _matmul - int8_static.int8_matmul_dynamic = _dynamic_matmul - int8_static.get_int8_matmul_backward_inputs = lambda input, hadamard: _quantize_int_mm(input.flatten(0, -2)) - if upstream_fixed: - int8_static.int8_matmul_backward_ckpt = lambda grad_output, input, weight, input_scale, scale, input_shape=None: None - else: - int8_static.int8_matmul_backward_ckpt = lambda grad_output, input, weight, input_scale, scale: None - - int8_dynamic = sys.modules["sdnq.training.layers.linear.linear_int8.linear_int8_dynamic_ckpt"] - int8_dynamic.quantize_int_mm = _quantize_int_mm - int8_dynamic.int8_matmul = _matmul - int8_dynamic.int8_matmul_dynamic = _dynamic_matmul - int8_dynamic.get_int8_matmul_dynamic_backward_inputs = lambda input, weight, hadamard: ( - *_quantize_int_mm(input.flatten(0, -2)), - *_quantize_int_mm(weight), - ) - int8_dynamic.int8_matmul_dynamic_backward_ckpt = lambda grad_output, input, weight, input_scale, weight_scale: None - - uint8_static = sys.modules["sdnq.training.layers.linear.linear_uint8.linear_uint8_ckpt"] - uint8_static.uint8_matmul = _matmul - uint8_static.uint8_matmul_dynamic = _dynamic_matmul - uint8_static.get_uint8_matmul_backward_inputs = lambda input, hadamard: _quantize_uint_mm(input.flatten(0, -2)) - uint8_static.uint8_matmul_backward_ckpt = ( - lambda grad_output, input, weight, input_scale, scale, input_zero_point, zero_point: None - ) - - uint8_dynamic = sys.modules["sdnq.training.layers.linear.linear_uint8.linear_uint8_dynamic_ckpt"] - uint8_dynamic.quantize_uint_mm = _quantize_uint_mm - uint8_dynamic.uint8_matmul = _matmul - uint8_dynamic.uint8_matmul_dynamic = _dynamic_matmul - uint8_dynamic.get_uint8_matmul_dynamic_backward_inputs = lambda input, weight, hadamard: ( - *_quantize_uint_mm(input.flatten(0, -2)), - *_quantize_uint_mm(weight), - ) - uint8_dynamic.uint8_matmul_dynamic_backward_ckpt = ( - lambda grad_output, input, weight, input_scale, weight_scale, input_zero_point, weight_zero_point: None - ) - - for dtype_name in ("fp8", "fp16"): - static = sys.modules[f"sdnq.training.layers.linear.linear_{dtype_name}.linear_{dtype_name}_ckpt"] - setattr(static, f"{dtype_name}_matmul", _matmul) - setattr(static, f"{dtype_name}_matmul_dynamic", _dynamic_matmul) - setattr(static, f"{dtype_name}_matmul_backward_ckpt", lambda grad_output, input, weight, input_scale, scale: None) - static.quantize_fp_mm = _quantize_fp_mm - - dynamic = sys.modules[f"sdnq.training.layers.linear.linear_{dtype_name}.linear_{dtype_name}_dynamic_ckpt"] - setattr(dynamic, f"{dtype_name}_matmul", _matmul) - setattr(dynamic, f"{dtype_name}_matmul_dynamic", _dynamic_matmul) - setattr( - dynamic, - f"get_{dtype_name}_matmul_dynamic_backward_inputs", - lambda input, weight, hadamard: (*_quantize_fp_mm(input.flatten(0, -2)), *_quantize_fp_mm(weight)), - ) - setattr( - dynamic, - f"{dtype_name}_matmul_dynamic_backward_ckpt", - lambda grad_output, input, weight, input_scale, weight_scale: None, - ) - dynamic.quantize_fp_mm = _quantize_fp_mm - - -class SDNQCompatTests(unittest.TestCase): - def setUp(self) -> None: - self.original_modules = {name: sys.modules.get(name) for name in MODULE_NAMES} - - def tearDown(self) -> None: - for name in list(sys.modules): - if name == "sdnq" or name.startswith("sdnq."): - sys.modules.pop(name) - for name, module in self.original_modules.items(): - if module is not None: - sys.modules[name] = module - - def test_version_gate_skips_old_sdnq(self): - _install_fake_sdnq_modules() - with patch("simpletuner.helpers.training.sdnq_compat.metadata.version", return_value="0.2.1"): - self.assertFalse(apply_sdnq_checkpointed_backward_fix()) - - def test_upstream_fixed_sdnq_is_left_unpatched(self): - _install_fake_sdnq_modules(upstream_fixed=True) - int8_static = sys.modules["sdnq.training.layers.linear.linear_int8.linear_int8_ckpt"] - original = int8_static.int8_matmul_backward_ckpt - with patch("simpletuner.helpers.training.sdnq_compat.metadata.version", return_value="0.2.2"): - self.assertFalse(apply_sdnq_checkpointed_backward_fix()) - self.assertIs(int8_static.int8_matmul_backward_ckpt, original) - - def test_patch_is_idempotent_and_updates_checkpoint_signatures(self): - _install_fake_sdnq_modules() - with patch("simpletuner.helpers.training.sdnq_compat.metadata.version", return_value="0.2.2"): - self.assertTrue(apply_sdnq_checkpointed_backward_fix()) - self.assertFalse(apply_sdnq_checkpointed_backward_fix()) - - function_names = { - "sdnq.training.layers.linear.linear_int8.linear_int8_ckpt": "int8_matmul_backward_ckpt", - "sdnq.training.layers.linear.linear_int8.linear_int8_dynamic_ckpt": "int8_matmul_dynamic_backward_ckpt", - "sdnq.training.layers.linear.linear_uint8.linear_uint8_ckpt": "uint8_matmul_backward_ckpt", - "sdnq.training.layers.linear.linear_uint8.linear_uint8_dynamic_ckpt": "uint8_matmul_dynamic_backward_ckpt", - "sdnq.training.layers.linear.linear_fp8.linear_fp8_ckpt": "fp8_matmul_backward_ckpt", - "sdnq.training.layers.linear.linear_fp8.linear_fp8_dynamic_ckpt": "fp8_matmul_dynamic_backward_ckpt", - "sdnq.training.layers.linear.linear_fp16.linear_fp16_ckpt": "fp16_matmul_backward_ckpt", - "sdnq.training.layers.linear.linear_fp16.linear_fp16_dynamic_ckpt": "fp16_matmul_dynamic_backward_ckpt", - } - for module_name, func_name in function_names.items(): - module = sys.modules[module_name] - self.assertIn("input_shape", signature(getattr(module, func_name)).parameters) - - def test_frozen_static_int8_weight_skips_backward_input_quantization(self): - _install_fake_sdnq_modules() - with patch("simpletuner.helpers.training.sdnq_compat.metadata.version", return_value="0.2.2"): - self.assertTrue(apply_sdnq_checkpointed_backward_fix()) - - int8_static = sys.modules["sdnq.training.layers.linear.linear_int8.linear_int8_ckpt"] - calls = {"backward_inputs": 0} - - def fail_if_called(input, hadamard): - calls["backward_inputs"] += 1 - raise AssertionError("input quantization should be skipped for frozen weights") - - int8_static.get_int8_matmul_backward_inputs = fail_if_called - - weight = torch.randn(4, 8) - weight.weight = weight - weight.scale = torch.ones(4, 1) - weight.zero_point = None - weight.svd_up = None - weight.svd_down = None - weight.sdnq_dequantizer = types.SimpleNamespace(use_hadamard=False, hadamard_group_size=128) - - input = torch.randn(2, 8, requires_grad=True) - output = int8_static.INT8MatmulBackwardCKPT.apply(input, weight, None) - output.sum().backward() - - self.assertEqual(0, calls["backward_inputs"]) - self.assertEqual(input.shape, input.grad.shape) - - -if __name__ == "__main__": - unittest.main()