SDNQ de-couple matmul dtype and add fp16 matmul

This commit is contained in:
Disty0
2025-11-22 02:16:20 +03:00
parent 9e91ac6310
commit b6e9332cfe
17 changed files with 451 additions and 162 deletions
+1
View File
@@ -199,6 +199,7 @@ def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.G
self = sdnq_quantize_layer(
self,
weights_dtype=sdnq_dequantizer.weights_dtype,
quantized_matmul_dtype=sdnq_dequantizer.quantized_matmul_dtype,
torch_dtype=sdnq_dequantizer.result_dtype,
group_size=sdnq_dequantizer.group_size,
svd_rank=sdnq_dequantizer.svd_rank,
+26 -5
View File
@@ -162,7 +162,7 @@ def get_sdnq_devices(mode="pre"):
return quantization_device, return_device
def create_sdnq_config(kwargs = None, allow: bool = True, module: str = 'Model', weights_dtype: str = None, modules_to_not_convert: list = None, modules_dtype_dict: dict = None):
def create_sdnq_config(kwargs = None, allow: bool = True, module: str = 'Model', weights_dtype: str = None, quantized_matmul_dtype: str = None, modules_to_not_convert: list = None, modules_dtype_dict: dict = None):
from modules import shared
if allow and (shared.opts.sdnq_quantize_mode in {'pre', 'auto'}) and (module == 'any' or module in shared.opts.sdnq_quantize_weights):
from modules.sdnq import SDNQConfig
@@ -175,6 +175,14 @@ def create_sdnq_config(kwargs = None, allow: bool = True, module: str = 'Model',
if weights_dtype is None or weights_dtype == 'none':
return kwargs
if quantized_matmul_dtype is None:
if module in {"TE", "LLM"} and shared.opts.sdnq_quantize_matmul_mode_te not in {"Same as model", "default"}:
quantized_matmul_dtype = shared.opts.sdnq_quantize_matmul_mode_te
else:
quantized_matmul_dtype = shared.opts.sdnq_quantize_matmul_mode
if quantized_matmul_dtype == "auto":
quantized_matmul_dtype = None
if modules_to_not_convert is None:
modules_to_not_convert = []
if modules_dtype_dict is None:
@@ -204,6 +212,7 @@ def create_sdnq_config(kwargs = None, allow: bool = True, module: str = 'Model',
sdnq_config = SDNQConfig(
weights_dtype=weights_dtype,
quantized_matmul_dtype=quantized_matmul_dtype,
group_size=shared.opts.sdnq_quantize_weights_group_size,
svd_rank=shared.opts.sdnq_svd_rank,
svd_steps=shared.opts.sdnq_svd_steps,
@@ -218,7 +227,9 @@ def create_sdnq_config(kwargs = None, allow: bool = True, module: str = 'Model',
modules_to_not_convert=modules_to_not_convert,
modules_dtype_dict=modules_dtype_dict.copy(),
)
log.debug(f'Quantization: module="{module}" type=sdnq mode=pre dtype={weights_dtype} matmul={shared.opts.sdnq_use_quantized_matmul} group_size={shared.opts.sdnq_quantize_weights_group_size} svd_rank={shared.opts.sdnq_svd_rank} svd_steps={shared.opts.sdnq_svd_steps} use_svd={shared.opts.sdnq_use_svd} quant_conv={shared.opts.sdnq_quantize_conv_layers} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} dequantize_fp32={shared.opts.sdnq_dequantize_fp32} quantize_with_gpu={shared.opts.sdnq_quantize_with_gpu} quantization_device={quantization_device} return_device={return_device} device_map={shared.opts.device_map} offload_mode={shared.opts.diffusers_offload_mode} non_blocking={shared.opts.diffusers_offload_nonblocking} modules_to_not_convert={modules_to_not_convert} modules_dtype_dict={modules_dtype_dict}')
if quantized_matmul_dtype is None:
quantized_matmul_dtype = "auto" # set for logging
log.debug(f'Quantization: module="{module}" type=sdnq mode=pre dtype={weights_dtype} matmul_dtype={quantized_matmul_dtype} matmul={shared.opts.sdnq_use_quantized_matmul} group_size={shared.opts.sdnq_quantize_weights_group_size} svd_rank={shared.opts.sdnq_svd_rank} svd_steps={shared.opts.sdnq_svd_steps} use_svd={shared.opts.sdnq_use_svd} quant_conv={shared.opts.sdnq_quantize_conv_layers} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} dequantize_fp32={shared.opts.sdnq_dequantize_fp32} quantize_with_gpu={shared.opts.sdnq_quantize_with_gpu} quantization_device={quantization_device} return_device={return_device} device_map={shared.opts.device_map} offload_mode={shared.opts.diffusers_offload_mode} non_blocking={shared.opts.diffusers_offload_nonblocking} modules_to_not_convert={modules_to_not_convert} modules_dtype_dict={modules_dtype_dict}')
if kwargs is None:
return sdnq_config
else:
@@ -477,7 +488,7 @@ def apply_layerwise(sd_model, quiet:bool=False):
log.error(f'Quantization: type=layerwise {e}')
def sdnq_quantize_model(model, op=None, sd_model=None, do_gc: bool = True, weights_dtype: str = None, modules_to_not_convert: list = None, modules_dtype_dict: dict = None):
def sdnq_quantize_model(model, op=None, sd_model=None, do_gc: bool = True, weights_dtype: str = None, quantized_matmul_dtype: str = None, modules_to_not_convert: list = None, modules_dtype_dict: dict = None):
global quant_last_model_name, quant_last_model_device # pylint: disable=global-statement
from modules import devices, shared, timer
from modules.sdnq import sdnq_post_load_quant
@@ -487,10 +498,17 @@ def sdnq_quantize_model(model, op=None, sd_model=None, do_gc: bool = True, weigh
weights_dtype = shared.opts.sdnq_quantize_weights_mode_te
else:
weights_dtype = shared.opts.sdnq_quantize_weights_mode
if weights_dtype is None or weights_dtype == 'none':
return model
if quantized_matmul_dtype is None:
if (op is not None) and ("text_encoder" in op or op in {"TE", "LLM"}) and (shared.opts.sdnq_quantize_matmul_mode_te not in {"Same as model", "default"}):
quantized_matmul_dtype = shared.opts.sdnq_quantize_matmul_mode_te
else:
quantized_matmul_dtype = shared.opts.sdnq_quantize_matmul_mode
if quantized_matmul_dtype == "auto":
quantized_matmul_dtype = None
quantization_device, return_device = get_sdnq_devices(mode="post")
if modules_to_not_convert is None:
@@ -523,6 +541,7 @@ def sdnq_quantize_model(model, op=None, sd_model=None, do_gc: bool = True, weigh
model = sdnq_post_load_quant(
model,
weights_dtype=weights_dtype,
quantized_matmul_dtype=quantized_matmul_dtype,
torch_dtype=devices.dtype,
group_size=shared.opts.sdnq_quantize_weights_group_size,
svd_rank=shared.opts.sdnq_svd_rank,
@@ -563,7 +582,9 @@ def sdnq_quantize_model(model, op=None, sd_model=None, do_gc: bool = True, weigh
if do_gc:
devices.torch_gc(force=True, reason='sdnq')
log.debug(f'Quantization: module="{op if op is not None else model.__class__}" type=sdnq mode=post dtype={weights_dtype} matmul={shared.opts.sdnq_use_quantized_matmul} svd={shared.opts.sdnq_use_svd}:group={shared.opts.sdnq_quantize_weights_group_size}:rank={shared.opts.sdnq_svd_rank}:steps={shared.opts.sdnq_svd_steps} quant_conv={shared.opts.sdnq_quantize_conv_layers} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} fp32={shared.opts.sdnq_dequantize_fp32} gpu={shared.opts.sdnq_quantize_with_gpu} device={quantization_device} return={return_device} map={shared.opts.device_map} non_blocking={shared.opts.diffusers_offload_nonblocking} modules_skip={modules_to_not_convert} modules_dtype={modules_dtype_dict}')
if quantized_matmul_dtype is None:
quantized_matmul_dtype = "auto" # set for logging
log.debug(f'Quantization: module="{op if op is not None else model.__class__}" type=sdnq mode=post dtype={weights_dtype} matmul_dtype={quantized_matmul_dtype} matmul={shared.opts.sdnq_use_quantized_matmul} svd={shared.opts.sdnq_use_svd}:group={shared.opts.sdnq_quantize_weights_group_size}:rank={shared.opts.sdnq_svd_rank}:steps={shared.opts.sdnq_svd_steps} quant_conv={shared.opts.sdnq_quantize_conv_layers} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} fp32={shared.opts.sdnq_dequantize_fp32} gpu={shared.opts.sdnq_quantize_with_gpu} device={quantization_device} return={return_device} map={shared.opts.device_map} non_blocking={shared.opts.diffusers_offload_nonblocking} modules_skip={modules_to_not_convert} modules_dtype={modules_dtype_dict}')
return model
+22 -3
View File
@@ -33,18 +33,23 @@ dtype_dict = {
"float8_e5m2": {"min": -57344, "max": 57344, "num_bits": 8, "sign": 1, "exponent": 5, "mantissa": 2, "target_dtype": torch.float8_e5m2, "torch_dtype": torch.float8_e5m2, "storage_dtype": torch.float8_e5m2, "is_unsigned": False, "is_integer": False, "is_packed": False},
}
dtype_dict["fp8"] = dtype_dict["float8_e4m3fn"]
dtype_dict["bool"] = dtype_dict["uint1"]
if hasattr(torch, "float8_e4m3fnuz"):
dtype_dict["float8_e4m3fnuz"] = {"min": -240, "max": 240, "num_bits": 8, "sign": 1, "exponent": 4, "mantissa": 3, "target_dtype": "fp8", "torch_dtype": torch.float8_e4m3fnuz, "storage_dtype": torch.float8_e4m3fnuz, "is_unsigned": False, "is_integer": False, "is_packed": False}
if hasattr(torch, "float8_e5m2fnuz"):
dtype_dict["float8_e5m2fnuz"] = {"min": -57344, "max": 57344, "num_bits": 8, "sign": 1, "exponent": 5, "mantissa": 2, "target_dtype": "fp8", "torch_dtype": torch.float8_e5m2fnuz, "storage_dtype": torch.float8_e5m2fnuz, "is_unsigned": False, "is_integer": False, "is_packed": False}
dtype_dict["fp32"] = dtype_dict["float32"]
dtype_dict["bf16"] = dtype_dict["bfloat16"]
dtype_dict["fp16"] = dtype_dict["float16"]
dtype_dict["fp8"] = dtype_dict["float8_e4m3fn"]
dtype_dict["bool"] = dtype_dict["uint1"]
linear_types = {"Linear"}
conv_types = {"Conv1d", "Conv2d", "Conv3d"}
conv_transpose_types = {"ConvTranspose1d", "ConvTranspose2d", "ConvTranspose3d"}
allowed_types = set.union(linear_types, conv_types, conv_transpose_types)
accepted_weights = set(dtype_dict.keys())
accepted_weight_dtypes = set(dtype_dict.keys())
accepted_matmul_dtypes = {"int8", "fp8", "fp16", "float8_e4m3fnuz", "float16"}
use_torch_compile = shared.opts.sdnq_dequantize_compile # this setting requires a full restart of the webui to apply
is_rdna2 = bool(devices.backend == "rocm" and int(getattr(torch.cuda.get_device_properties(devices.device), "gcnArchName", "gfx0000")[3:]) < 1100)
@@ -77,6 +82,20 @@ else:
int_mm_func = torch._int_mm
fp_mm_func = None
if os.environ.get("SDNQ_USE_TRITON_MM", "1").lower() not in {"0", "false", "no"}:
try:
from .triton_mm import fp_mm
fp_mm_func = fp_mm
except ImportError:
fp_mm_func = None
if fp_mm_func is None:
def fp_mm(x: torch.Tensor, y: torch.Tensor) -> torch.FloatTensor:
return torch.mm(x,y, out_dtype=torch.float32)
fp_mm_func = fp_mm
if use_torch_compile:
torch._dynamo.config.cache_size_limit = max(8192, torch._dynamo.config.cache_size_limit)
torch._dynamo.config.accumulated_recompile_limit = max(8192, torch._dynamo.config.accumulated_recompile_limit)
+42 -44
View File
@@ -74,92 +74,87 @@ def dequantize_packed_int_symmetric(weight: torch.ByteTensor, scale: torch.Float
@devices.inference_context()
def quantize_int8(input: torch.FloatTensor, dim: int = -1) -> Tuple[torch.CharTensor, torch.FloatTensor]:
scale = torch.amax(input.abs(), dim=dim, keepdims=True).div_(127)
input = torch.div(input, scale).round_().clamp_(-128, 127).to(dtype=torch.int8)
def quantize_int_mm(input: torch.FloatTensor, dim: int = -1, matmul_dtype: str = "int8") -> Tuple[torch.Tensor, torch.FloatTensor]:
scale = torch.amax(input.abs(), dim=dim, keepdims=True).div_(dtype_dict[matmul_dtype]["max"])
input = torch.div(input, scale).round_().clamp_(dtype_dict[matmul_dtype]["min"], dtype_dict[matmul_dtype]["max"]).to(dtype=dtype_dict[matmul_dtype]["torch_dtype"])
return input, scale
@devices.inference_context()
def quantize_int8_sr(input: torch.FloatTensor, dim: int = -1) -> Tuple[torch.CharTensor, torch.FloatTensor]:
scale = torch.amax(input.abs(), dim=dim, keepdims=True).div_(127)
def quantize_int_mm_sr(input: torch.FloatTensor, dim: int = -1, matmul_dtype: str = "int8") -> Tuple[torch.Tensor, torch.FloatTensor]:
scale = torch.amax(input.abs(), dim=dim, keepdims=True).div_(dtype_dict[matmul_dtype]["max"])
input = torch.normal(0, 0.1, input.shape, device=input.device, dtype=input.dtype
).addcdiv_(input, scale).round_().clamp_(-128, 127).to(dtype=torch.int8)
).addcdiv_(input, scale).round_().clamp_(dtype_dict[matmul_dtype]["min"], dtype_dict[matmul_dtype]["max"]).to(dtype=dtype_dict[matmul_dtype]["torch_dtype"])
return input, scale
@devices.inference_context()
def quantize_fp8(input: torch.FloatTensor, dim: int = -1, is_e5: bool = False) -> Tuple[torch.Tensor, torch.FloatTensor]:
if is_e5:
max_range = 57344
fp8_dtype = torch.float8_e5m2
else:
max_range = 448
fp8_dtype = torch.float8_e4m3fn
scale = torch.amax(input.abs(), dim=dim, keepdims=True).div_(max_range)
input = torch.div(input, scale).nan_to_num_().clamp_(-max_range, max_range).to(dtype=fp8_dtype)
def quantize_fp_mm(input: torch.FloatTensor, dim: int = -1, matmul_dtype: str = "float8_e4m3fn") -> Tuple[torch.Tensor, torch.FloatTensor]:
scale = torch.amax(input.abs(), dim=dim, keepdims=True).div_(dtype_dict[matmul_dtype]["max"])
input = torch.div(input, scale).nan_to_num_().clamp_(dtype_dict[matmul_dtype]["min"], dtype_dict[matmul_dtype]["max"]).to(dtype=dtype_dict[matmul_dtype]["torch_dtype"])
return input, scale
@devices.inference_context()
def quantize_fp8_sr(input: torch.FloatTensor, dim: int = -1, is_e5: bool = False) -> Tuple[torch.Tensor, torch.FloatTensor]:
if is_e5:
max_range = 57344
fp8_dtype = torch.float8_e5m2
mantissa_difference = 2097152
else:
max_range = 448
fp8_dtype = torch.float8_e4m3fn
mantissa_difference = 1048576
scale = torch.amax(input.abs(), dim=dim, keepdims=True).div_(max_range)
def quantize_fp_mm_sr(input: torch.FloatTensor, dim: int = -1, matmul_dtype: str = "float8_e4m3fn") -> Tuple[torch.Tensor, torch.FloatTensor]:
mantissa_difference = mantissa_difference = 1 << (23 - dtype_dict[matmul_dtype]["mantissa"])
scale = torch.amax(input.abs(), dim=dim, keepdims=True).div_(dtype_dict[matmul_dtype]["max"])
input = torch.div(input, scale).to(dtype=torch.float32).view(dtype=torch.int32)
input = input.add_(torch.randint_like(input, low=0, high=mantissa_difference)).bitwise_and_(-mantissa_difference).view(dtype=torch.float32)
input = input.nan_to_num_().clamp_(-max_range, max_range).to(dtype=fp8_dtype)
input = input.nan_to_num_().clamp_(dtype_dict[matmul_dtype]["min"], dtype_dict[matmul_dtype]["max"]).to(dtype=dtype_dict[matmul_dtype]["torch_dtype"])
return input, scale
@devices.inference_context()
def re_quantize_int8(weight: torch.FloatTensor) -> Tuple[torch.CharTensor, torch.FloatTensor]:
def re_quantize_int_mm(weight: torch.FloatTensor) -> Tuple[torch.Tensor, torch.FloatTensor]:
if weight.ndim > 2: # convs
weight = weight.flatten(1,-1)
if use_contiguous_mm:
weight, scale = quantize_int8(weight.t(), dim=-0)
weight, scale = quantize_int_mm(weight.t(), dim=-0)
weight, scale = weight.contiguous(), scale.contiguous()
else:
weight, scale = quantize_int8(weight.contiguous(), dim=-1)
weight, scale = quantize_int_mm(weight.contiguous(), dim=-1)
weight, scale = weight.t_(), scale.t_()
return weight, scale
@devices.inference_context()
def re_quantize_fp8(weight: torch.FloatTensor, is_e5: bool = False) -> Tuple[torch.CharTensor, torch.FloatTensor]:
def re_quantize_fp_mm(weight: torch.FloatTensor, matmul_dtype: str = "float8_e4m3fn") -> Tuple[torch.Tensor, torch.FloatTensor]:
if weight.ndim > 2: # convs
weight = weight.flatten(1,-1)
weight, scale = quantize_fp8(weight.contiguous(), dim=-1, is_e5=is_e5)
weight, scale = quantize_fp_mm(weight.contiguous(), dim=-1, matmul_dtype=matmul_dtype)
weight, scale = weight.t_(), scale.t_()
if not use_tensorwise_fp8_matmul:
if not use_tensorwise_fp8_matmul and dtype_dict[matmul_dtype]["num_bits"] == 8:
scale = scale.to(dtype=torch.float32)
return weight, scale
@devices.inference_context()
def re_quantize_matmul_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, result_shape: Optional[torch.Size] = None, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None) -> Tuple[torch.CharTensor, torch.FloatTensor]:
return re_quantize_int8(dequantize_asymmetric(weight, scale, zero_point, svd_up=svd_up, svd_down=svd_down, dtype=scale.dtype, result_shape=result_shape))
def re_quantize_matmul_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, matmul_dtype: str, result_shape: Optional[torch.Size] = None, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None) -> Tuple[torch.Tensor, torch.FloatTensor]:
weight = dequantize_asymmetric(weight, scale, zero_point, svd_up=svd_up, svd_down=svd_down, dtype=scale.dtype, result_shape=result_shape)
if dtype_dict[matmul_dtype]["is_integer"]:
return re_quantize_int_mm(weight)
else:
return re_quantize_fp_mm(weight, matmul_dtype=matmul_dtype)
@devices.inference_context()
def re_quantize_matmul_symmetric(weight: torch.CharTensor, scale: torch.FloatTensor, result_shape: Optional[torch.Size] = None, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None) -> Tuple[torch.CharTensor, torch.FloatTensor]:
return re_quantize_int8(dequantize_symmetric(weight, scale, svd_up=svd_up, svd_down=svd_down, dtype=scale.dtype, result_shape=result_shape))
def re_quantize_matmul_symmetric(weight: torch.CharTensor, scale: torch.FloatTensor, matmul_dtype: str, result_shape: Optional[torch.Size] = None, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None) -> Tuple[torch.Tensor, torch.FloatTensor]:
weight = dequantize_symmetric(weight, scale, svd_up=svd_up, svd_down=svd_down, dtype=scale.dtype, result_shape=result_shape)
if dtype_dict[matmul_dtype]["is_integer"]:
return re_quantize_int_mm(weight)
else:
return re_quantize_fp_mm(weight, matmul_dtype=matmul_dtype)
@devices.inference_context()
def re_quantize_matmul_packed_int_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, weights_dtype: str, result_shape: torch.Size, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None) -> Tuple[torch.CharTensor, torch.FloatTensor]:
return re_quantize_matmul_asymmetric(unpack_int_asymetric(weight, shape, weights_dtype), scale, zero_point, svd_up=svd_up, svd_down=svd_down, result_shape=result_shape)
def re_quantize_matmul_packed_int_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, weights_dtype: str, matmul_dtype: str, result_shape: torch.Size, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None) -> Tuple[torch.Tensor, torch.FloatTensor]:
return re_quantize_matmul_asymmetric(unpack_int_asymetric(weight, shape, weights_dtype), scale, zero_point, matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=result_shape)
@devices.inference_context()
def re_quantize_matmul_packed_int_symmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, result_shape: Optional[torch.Size] = None, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None) -> Tuple[torch.CharTensor, torch.FloatTensor]:
return re_quantize_matmul_symmetric(unpack_int_symetric(weight, shape, weights_dtype, dtype=scale.dtype), scale, svd_up=svd_up, svd_down=svd_down, result_shape=result_shape)
def re_quantize_matmul_packed_int_symmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, matmul_dtype: str, result_shape: Optional[torch.Size] = None, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None) -> Tuple[torch.Tensor, torch.FloatTensor]:
return re_quantize_matmul_symmetric(unpack_int_symetric(weight, shape, weights_dtype, dtype=scale.dtype), scale, matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=result_shape)
@devices.inference_context()
@@ -221,6 +216,7 @@ class SDNQDequantizer():
original_stride: List[int],
quantized_weight_shape: torch.Size,
weights_dtype: str,
quantized_matmul_dtype: str,
group_size: int,
svd_rank: int,
svd_steps: int,
@@ -232,12 +228,14 @@ class SDNQDequantizer():
self.is_packed = dtype_dict[weights_dtype]["is_packed"]
self.is_unsigned = dtype_dict[weights_dtype]["is_unsigned"]
self.is_integer = dtype_dict[weights_dtype]["is_integer"]
self.is_integer_matmul = dtype_dict[quantized_matmul_dtype]["is_integer"]
self.result_dtype = result_dtype
self.result_shape = result_shape
self.original_shape = original_shape
self.original_stride = original_stride
self.quantized_weight_shape = quantized_weight_shape
self.weights_dtype = weights_dtype
self.quantized_matmul_dtype = quantized_matmul_dtype
self.group_size = group_size
self.svd_rank = svd_rank
self.svd_steps = svd_steps
@@ -250,14 +248,14 @@ class SDNQDequantizer():
def re_quantize_matmul(self, weight, scale, zero_point, svd_up, svd_down): # pylint: disable=unused-argument
if self.is_packed:
if self.is_unsigned:
return re_quantize_matmul_packed_int_asymmetric_compiled(weight, scale, zero_point, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape)
return re_quantize_matmul_packed_int_asymmetric_compiled(weight, scale, zero_point, self.quantized_weight_shape, self.weights_dtype, self.quantized_matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape)
else:
return re_quantize_matmul_packed_int_symmetric_compiled(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape)
return re_quantize_matmul_packed_int_symmetric_compiled(weight, scale, self.quantized_weight_shape, self.weights_dtype, self.quantized_matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape)
else:
if self.is_unsigned:
return re_quantize_matmul_asymmetric_compiled(weight, scale, zero_point, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape)
return re_quantize_matmul_asymmetric_compiled(weight, scale, zero_point, self.quantized_matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape)
else:
return re_quantize_matmul_symmetric_compiled(weight, scale, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape)
return re_quantize_matmul_symmetric_compiled(weight, scale, self.quantized_matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape)
@devices.inference_context()
def __call__(self, weight, scale, zero_point, svd_up, svd_down, skip_quantized_matmul: bool = False, dtype: torch.dtype = None): # pylint: disable=unused-argument
+22 -14
View File
@@ -2,22 +2,26 @@
from typing import Callable
from .common import conv_types, conv_transpose_types, use_tensorwise_fp8_matmul
from .common import dtype_dict, conv_types, conv_transpose_types, use_tensorwise_fp8_matmul
def get_forward_func(layer_class_name: str, use_quantized_matmul: bool, is_integer: bool) -> Callable: # pylint: disable=inconsistent-return-statements
def get_forward_func(layer_class_name: str, quantized_matmul_dtype: str, use_quantized_matmul: bool) -> Callable: # pylint: disable=inconsistent-return-statements
if layer_class_name in conv_types:
if use_quantized_matmul:
if is_integer:
if dtype_dict[quantized_matmul_dtype]["is_integer"]:
from .layers.conv.conv_int8 import quantized_conv_forward_int8_matmul
return quantized_conv_forward_int8_matmul
else:
if use_tensorwise_fp8_matmul:
from .layers.conv.conv_fp8_tensorwise import quantized_conv_forward_fp8_matmul_tensorwise
return quantized_conv_forward_fp8_matmul_tensorwise
if dtype_dict[quantized_matmul_dtype]["num_bits"] == 8:
if use_tensorwise_fp8_matmul:
from .layers.conv.conv_fp8_tensorwise import quantized_conv_forward_fp8_matmul_tensorwise
return quantized_conv_forward_fp8_matmul_tensorwise
else:
from .layers.conv.conv_fp8 import quantized_conv_forward_fp8_matmul
return quantized_conv_forward_fp8_matmul
else:
from .layers.conv.conv_fp8 import quantized_conv_forward_fp8_matmul
return quantized_conv_forward_fp8_matmul
from .layers.conv.conv_fp16 import quantized_conv_forward_fp16_matmul
return quantized_conv_forward_fp16_matmul
else:
from .layers.conv.forward import quantized_conv_forward
return quantized_conv_forward
@@ -33,16 +37,20 @@ def get_forward_func(layer_class_name: str, use_quantized_matmul: bool, is_integ
return quantized_conv_transpose_3d_forward
else:
if use_quantized_matmul:
if is_integer:
if dtype_dict[quantized_matmul_dtype]["is_integer"]:
from .layers.linear.linear_int8 import quantized_linear_forward_int8_matmul
return quantized_linear_forward_int8_matmul
else:
if use_tensorwise_fp8_matmul:
from .layers.linear.linear_fp8_tensorwise import quantized_linear_forward_fp8_matmul_tensorwise
return quantized_linear_forward_fp8_matmul_tensorwise
if dtype_dict[quantized_matmul_dtype]["num_bits"] == 8:
if use_tensorwise_fp8_matmul:
from .layers.linear.linear_fp8_tensorwise import quantized_linear_forward_fp8_matmul_tensorwise
return quantized_linear_forward_fp8_matmul_tensorwise
else:
from .layers.linear.linear_fp8 import quantized_linear_forward_fp8_matmul
return quantized_linear_forward_fp8_matmul
else:
from .layers.linear.linear_fp8 import quantized_linear_forward_fp8_matmul
return quantized_linear_forward_fp8_matmul
from .layers.linear.linear_fp16 import quantized_linear_forward_fp16_matmul
return quantized_linear_forward_fp16_matmul
else:
from .layers.linear.forward import quantized_linear_forward
return quantized_linear_forward
+81
View File
@@ -0,0 +1,81 @@
# pylint: disable=relative-beyond-top-level,redefined-builtin,protected-access
from typing import List
import torch
from ...common import compile_func, fp_mm_func # noqa: TID252
from ...dequantizer import dequantize_symmetric, dequantize_symmetric_with_bias # noqa: TID252
from .forward import get_conv_args, process_conv_input
from ..linear.linear_fp8_tensorwise import quantize_fp_mm_input_tensorwise # noqa: TID252
from ..linear.forward import check_mats # noqa: TID252
def conv_fp16_matmul(
input: torch.FloatTensor,
weight: torch.Tensor,
scale: torch.FloatTensor,
result_shape: torch.Size,
reversed_padding_repeated_twice: List[int],
padding_mode: str, conv_type: int,
groups: int, stride: List[int],
padding: List[int], dilation: List[int],
bias: torch.FloatTensor = None,
svd_up: torch.FloatTensor = None,
svd_down: torch.FloatTensor = None,
) -> torch.FloatTensor:
return_dtype = input.dtype
input, mm_output_shape = process_conv_input(conv_type, input, reversed_padding_repeated_twice, padding_mode, result_shape, stride, padding, dilation)
if svd_up is not None:
input = input.flatten(0,-2)
if bias is not None:
bias = torch.addmm(bias.to(dtype=svd_down.dtype), torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
else:
bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
input, scale = quantize_fp_mm_input_tensorwise(input, scale, matmul_dtype="float16")
input, weight = check_mats(input, weight)
if groups == 1:
result = fp_mm_func(input, weight)
else:
weight = weight.view(weight.shape[0], groups, weight.shape[1] // groups)
input = input.view(input.shape[0], groups, input.shape[1] // groups)
result = []
for i in range(groups):
result.append(fp_mm_func(input[:, i], weight[:, i]))
result = torch.cat(result, dim=-1)
if bias is not None:
dequantize_symmetric_with_bias(result, scale, bias, dtype=return_dtype, result_shape=mm_output_shape)
else:
dequantize_symmetric(result, scale, dtype=return_dtype, result_shape=mm_output_shape)
if conv_type == 1:
result = result.transpose_(1,2)
elif conv_type == 2:
result = result.permute(0,3,1,2)
elif conv_type == 3:
result = result.permute(0,4,1,2,3)
return result
def quantized_conv_forward_fp16_matmul(self, input) -> torch.FloatTensor:
if self.sdnq_dequantizer.re_quantize_for_matmul:
weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None)
else:
weight, scale = self.weight, self.scale
conv_type, stride, padding, dilation = get_conv_args(input.ndim, self.stride, self.padding, self.dilation)
return conv_fp16_matmul(
input, weight, scale,
self.sdnq_dequantizer.result_shape,
self._reversed_padding_repeated_twice,
self.padding_mode, conv_type,
self.groups, stride, padding, dilation,
bias=self.bias,
svd_up=self.svd_up,
svd_down=self.svd_down,
)
conv_fp16_matmul = compile_func(conv_fp16_matmul)
+15 -8
View File
@@ -5,23 +5,24 @@ from typing import List
import torch
from ...common import compile_func # noqa: TID252
from ..linear.linear_fp8 import quantize_fp8_matmul_input # noqa: TID252
from ..linear.forward import check_mats # noqa: TID252
from .forward import get_conv_args, process_conv_input
from ..linear.linear_fp8 import quantize_fp_mm_input # noqa: TID252
from ..linear.forward import check_mats # noqa: TID252
def conv_fp8_matmul(
input: torch.FloatTensor,
weight: torch.Tensor,
bias: torch.FloatTensor,
scale: torch.FloatTensor,
svd_up: torch.FloatTensor,
svd_down: torch.FloatTensor,
result_shape: torch.Size,
reversed_padding_repeated_twice: List[int],
padding_mode: str, conv_type: int,
groups: int, stride: List[int],
padding: List[int], dilation: List[int],
bias: torch.FloatTensor = None,
svd_up: torch.FloatTensor = None,
svd_down: torch.FloatTensor = None,
) -> torch.FloatTensor:
return_dtype = input.dtype
input, mm_output_shape = process_conv_input(conv_type, input, reversed_padding_repeated_twice, padding_mode, result_shape, stride, padding, dilation)
@@ -29,7 +30,7 @@ def conv_fp8_matmul(
input = input.flatten(0,-2)
svd_bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
input, input_scale = quantize_fp8_matmul_input(input)
input, input_scale = quantize_fp_mm_input(input)
input, weight = check_mats(input, weight)
if groups == 1:
@@ -68,14 +69,20 @@ def conv_fp8_matmul(
def quantized_conv_forward_fp8_matmul(self, input) -> torch.FloatTensor:
if torch.numel(input) / input.shape[2] < 32:
return self._conv_forward(input, self.sdnq_dequantizer(self.weight, self.scale, self.zero_point, self.svd_up, self.svd_down, skip_quantized_matmul=True), self.bias)
if self.sdnq_dequantizer.re_quantize_for_matmul:
weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None)
else:
weight, scale = self.weight, self.scale
conv_type, stride, padding, dilation = get_conv_args(input.ndim, self.stride, self.padding, self.dilation)
return conv_fp8_matmul(
input, self.weight, self.bias,
self.scale, self.svd_up, self.svd_down,
input, weight, scale,
self.sdnq_dequantizer.result_shape,
self._reversed_padding_repeated_twice,
self.padding_mode, conv_type,
self.groups, stride, padding, dilation,
bias=self.bias,
svd_up=self.svd_up,
svd_down=self.svd_down,
)
@@ -6,23 +6,24 @@ import torch
from ...common import compile_func # noqa: TID252
from ...dequantizer import dequantize_symmetric, dequantize_symmetric_with_bias # noqa: TID252
from ..linear.linear_fp8_tensorwise import quantize_fp8_matmul_input_tensorwise # noqa: TID252
from ..linear.forward import check_mats # noqa: TID252
from .forward import get_conv_args, process_conv_input
from ..linear.linear_fp8_tensorwise import quantize_fp_mm_input_tensorwise # noqa: TID252
from ..linear.forward import check_mats # noqa: TID252
def conv_fp8_matmul_tensorwise(
input: torch.FloatTensor,
weight: torch.Tensor,
bias: torch.FloatTensor,
scale: torch.FloatTensor,
svd_up: torch.FloatTensor,
svd_down: torch.FloatTensor,
result_shape: torch.Size,
reversed_padding_repeated_twice: List[int],
padding_mode: str, conv_type: int,
groups: int, stride: List[int],
padding: List[int], dilation: List[int],
bias: torch.FloatTensor = None,
svd_up: torch.FloatTensor = None,
svd_down: torch.FloatTensor = None,
) -> torch.FloatTensor:
return_dtype = input.dtype
input, mm_output_shape = process_conv_input(conv_type, input, reversed_padding_repeated_twice, padding_mode, result_shape, stride, padding, dilation)
@@ -33,7 +34,7 @@ def conv_fp8_matmul_tensorwise(
else:
bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
input, scale = quantize_fp8_matmul_input_tensorwise(input, scale)
input, scale = quantize_fp_mm_input_tensorwise(input, scale)
input, weight = check_mats(input, weight)
dummy_input_scale = torch.ones(1, device=input.device, dtype=torch.float32)
@@ -63,14 +64,20 @@ def conv_fp8_matmul_tensorwise(
def quantized_conv_forward_fp8_matmul_tensorwise(self, input) -> torch.FloatTensor:
if torch.numel(input) / input.shape[2] < 32:
return self._conv_forward(input, self.sdnq_dequantizer(self.weight, self.scale, self.zero_point, self.svd_up, self.svd_down, skip_quantized_matmul=True), self.bias)
if self.sdnq_dequantizer.re_quantize_for_matmul:
weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None)
else:
weight, scale = self.weight, self.scale
conv_type, stride, padding, dilation = get_conv_args(input.ndim, self.stride, self.padding, self.dilation)
return conv_fp8_matmul_tensorwise(
input, self.weight, self.bias,
self.scale, self.svd_up, self.svd_down,
input, weight, scale,
self.sdnq_dequantizer.result_shape,
self._reversed_padding_repeated_twice,
self.padding_mode, conv_type,
self.groups, stride, padding, dilation,
bias=self.bias,
svd_up=self.svd_up,
svd_down=self.svd_down,
)
+4 -3
View File
@@ -7,9 +7,10 @@ import torch
from ...common import compile_func, int_mm_func # noqa: TID252
from ...packed_int import unpack_int_symetric # noqa: TID252
from ...dequantizer import dequantize_symmetric, dequantize_symmetric_with_bias # noqa: TID252
from ..linear.linear_int8 import quantize_int8_matmul_input # noqa: TID252
from ..linear.forward import check_mats # noqa: TID252
from .forward import get_conv_args, process_conv_input
from ..linear.linear_int8 import quantize_int_mm_input # noqa: TID252
from ..linear.forward import check_mats # noqa: TID252
def conv_int8_matmul(
@@ -36,7 +37,7 @@ def conv_int8_matmul(
else:
bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
input, scale = quantize_int8_matmul_input(input, scale)
input, scale = quantize_int_mm_input(input, scale)
if quantized_weight_shape is not None:
weight = unpack_int_symetric(weight, quantized_weight_shape, weights_dtype, dtype=torch.int8)
input, weight = check_mats(input, weight)
+44
View File
@@ -0,0 +1,44 @@
# pylint: disable=relative-beyond-top-level,redefined-builtin,protected-access
import torch
from ...common import compile_func, fp_mm_func # noqa: TID252
from ...dequantizer import dequantize_symmetric, dequantize_symmetric_with_bias # noqa: TID252
from .forward import check_mats
from .linear_fp8_tensorwise import quantize_fp_mm_input_tensorwise
def fp16_matmul(
input: torch.FloatTensor,
weight: torch.Tensor,
scale: torch.FloatTensor,
bias: torch.FloatTensor = None,
svd_up: torch.FloatTensor = None,
svd_down: torch.FloatTensor = None,
) -> torch.FloatTensor:
return_dtype = input.dtype
output_shape = (*input.shape[:-1], weight.shape[-1])
if svd_up is not None:
input.flatten(0,-2)
if bias is not None:
bias = torch.addmm(bias.to(dtype=svd_down.dtype), torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
else:
bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
input, scale = quantize_fp_mm_input_tensorwise(input, scale, matmul_dtype="float16")
input, weight = check_mats(input, weight)
if bias is not None:
return dequantize_symmetric_with_bias(fp_mm_func(input, weight), scale, bias, dtype=return_dtype, result_shape=output_shape)
else:
return dequantize_symmetric(fp_mm_func(input, weight), scale, dtype=return_dtype, result_shape=output_shape)
def quantized_linear_forward_fp16_matmul(self, input: torch.FloatTensor) -> torch.FloatTensor:
if self.sdnq_dequantizer.re_quantize_for_matmul:
weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None)
else:
weight, scale = self.weight, self.scale
return fp16_matmul(input, weight, scale, bias=self.bias, svd_up=self.svd_up, svd_down=self.svd_down)
fp16_matmul = compile_func(fp16_matmul)
+13 -8
View File
@@ -5,30 +5,31 @@ from typing import Tuple
import torch
from ...common import compile_func # noqa: TID252
from ...dequantizer import quantize_fp8 # noqa: TID252
from ...dequantizer import quantize_fp_mm # noqa: TID252
from .forward import check_mats
def quantize_fp8_matmul_input(input: torch.FloatTensor) -> Tuple[torch.Tensor, torch.FloatTensor]:
def quantize_fp_mm_input(input: torch.FloatTensor, matmul_dtype: str = "float8_e4m3fn") -> Tuple[torch.Tensor, torch.FloatTensor]:
input = input.flatten(0,-2).to(dtype=torch.float32)
input, input_scale = quantize_fp8(input, dim=-1)
input, input_scale = quantize_fp_mm(input, dim=-1, matmul_dtype=matmul_dtype)
return input, input_scale
def fp8_matmul(
input: torch.FloatTensor,
weight: torch.Tensor,
bias: torch.FloatTensor,
scale: torch.FloatTensor,
svd_up: torch.FloatTensor,
svd_down: torch.FloatTensor,
bias: torch.FloatTensor = None,
svd_up: torch.FloatTensor = None,
svd_down: torch.FloatTensor = None,
) -> torch.FloatTensor:
return_dtype = input.dtype
output_shape = (*input.shape[:-1], weight.shape[-1])
if svd_up is not None:
input = input.flatten(0,-2)
svd_bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
input, input_scale = quantize_fp8_matmul_input(input)
input, input_scale = quantize_fp_mm_input(input)
input, weight = check_mats(input, weight)
if bias is not None and bias.dtype != torch.bfloat16:
bias = bias.to(dtype=torch.bfloat16)
@@ -42,7 +43,11 @@ def fp8_matmul(
def quantized_linear_forward_fp8_matmul(self, input: torch.FloatTensor) -> torch.FloatTensor:
if torch.numel(input) / input.shape[-1] < 32:
return torch.nn.functional.linear(input, self.sdnq_dequantizer(self.weight, self.scale, self.zero_point, self.svd_up, self.svd_down, skip_quantized_matmul=True), self.bias)
return fp8_matmul(input, self.weight, self.bias, self.scale, self.svd_up, self.svd_down)
if self.sdnq_dequantizer.re_quantize_for_matmul:
weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None)
else:
weight, scale = self.weight, self.scale
return fp8_matmul(input, weight, scale, bias=self.bias, svd_up=self.svd_up, svd_down=self.svd_down)
fp8_matmul = compile_func(fp8_matmul)
@@ -5,13 +5,14 @@ from typing import Tuple
import torch
from ...common import compile_func # noqa: TID252
from ...dequantizer import quantize_fp8, dequantize_symmetric, dequantize_symmetric_with_bias # noqa: TID252
from ...dequantizer import quantize_fp_mm, dequantize_symmetric, dequantize_symmetric_with_bias # noqa: TID252
from .forward import check_mats
def quantize_fp8_matmul_input_tensorwise(input: torch.FloatTensor, scale: torch.FloatTensor) -> Tuple[torch.Tensor, torch.FloatTensor]:
def quantize_fp_mm_input_tensorwise(input: torch.FloatTensor, scale: torch.FloatTensor, matmul_dtype: str = "float8_e4m3fn") -> Tuple[torch.Tensor, torch.FloatTensor]:
input = input.flatten(0,-2).to(dtype=scale.dtype)
input, input_scale = quantize_fp8(input, dim=-1)
input, input_scale = quantize_fp_mm(input, dim=-1, matmul_dtype=matmul_dtype)
scale = torch.mul(input_scale, scale)
if scale.dtype == torch.float16: # fp16 will overflow
scale = scale.to(dtype=torch.float32)
@@ -21,10 +22,10 @@ def quantize_fp8_matmul_input_tensorwise(input: torch.FloatTensor, scale: torch.
def fp8_matmul_tensorwise(
input: torch.FloatTensor,
weight: torch.Tensor,
bias: torch.FloatTensor,
scale: torch.FloatTensor,
svd_up: torch.FloatTensor,
svd_down: torch.FloatTensor,
bias: torch.FloatTensor = None,
svd_up: torch.FloatTensor = None,
svd_down: torch.FloatTensor = None,
) -> torch.FloatTensor:
return_dtype = input.dtype
output_shape = (*input.shape[:-1], weight.shape[-1])
@@ -35,7 +36,7 @@ def fp8_matmul_tensorwise(
else:
bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
dummy_input_scale = torch.ones(1, device=input.device, dtype=torch.float32)
input, scale = quantize_fp8_matmul_input_tensorwise(input, scale)
input, scale = quantize_fp_mm_input_tensorwise(input, scale)
input, weight = check_mats(input, weight)
if bias is not None:
return dequantize_symmetric_with_bias(torch._scaled_mm(input, weight, scale_a=dummy_input_scale, scale_b=dummy_input_scale, bias=None, out_dtype=scale.dtype), scale, bias, dtype=return_dtype, result_shape=output_shape)
@@ -46,7 +47,11 @@ def fp8_matmul_tensorwise(
def quantized_linear_forward_fp8_matmul_tensorwise(self, input: torch.FloatTensor) -> torch.FloatTensor:
if torch.numel(input) / input.shape[-1] < 32:
return torch.nn.functional.linear(input, self.sdnq_dequantizer(self.weight, self.scale, self.zero_point, self.svd_up, self.svd_down, skip_quantized_matmul=True), self.bias)
return fp8_matmul_tensorwise(input, self.weight, self.bias, self.scale, self.svd_up, self.svd_down)
if self.sdnq_dequantizer.re_quantize_for_matmul:
weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None)
else:
weight, scale = self.weight, self.scale
return fp8_matmul_tensorwise(input, weight, scale, bias=self.bias, svd_up=self.svd_up, svd_down=self.svd_down)
fp8_matmul_tensorwise = compile_func(fp8_matmul_tensorwise)
+18 -10
View File
@@ -6,13 +6,14 @@ import torch
from ...common import compile_func, int_mm_func # noqa: TID252
from ...packed_int import unpack_int_symetric # noqa: TID252
from ...dequantizer import quantize_int8, dequantize_symmetric, dequantize_symmetric_with_bias # noqa: TID252
from ...dequantizer import quantize_int_mm, dequantize_symmetric, dequantize_symmetric_with_bias # noqa: TID252
from .forward import check_mats
def quantize_int8_matmul_input(input: torch.FloatTensor, scale: torch.FloatTensor) -> Tuple[torch.CharTensor, torch.FloatTensor]:
def quantize_int_mm_input(input: torch.FloatTensor, scale: torch.FloatTensor) -> Tuple[torch.CharTensor, torch.FloatTensor]:
input = input.flatten(0,-2).to(dtype=scale.dtype)
input, input_scale = quantize_int8(input, dim=-1)
input, input_scale = quantize_int_mm(input, dim=-1)
scale = torch.mul(input_scale, scale)
if scale.dtype == torch.float16: # fp16 will overflow
scale = scale.to(dtype=torch.float32)
@@ -22,12 +23,12 @@ def quantize_int8_matmul_input(input: torch.FloatTensor, scale: torch.FloatTenso
def int8_matmul(
input: torch.FloatTensor,
weight: torch.Tensor,
bias: torch.FloatTensor,
scale: torch.FloatTensor,
svd_up: torch.FloatTensor,
svd_down: torch.FloatTensor,
quantized_weight_shape: torch.Size,
weights_dtype: str,
bias: torch.FloatTensor = None,
svd_up: torch.FloatTensor = None,
svd_down: torch.FloatTensor = None,
quantized_weight_shape: torch.Size = None,
weights_dtype: str = None,
) -> torch.FloatTensor:
if quantized_weight_shape is not None:
weight = unpack_int_symetric(weight, quantized_weight_shape, weights_dtype, dtype=torch.int8)
@@ -39,7 +40,7 @@ def int8_matmul(
bias = torch.addmm(bias.to(dtype=svd_down.dtype), torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
else:
bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
input, scale = quantize_int8_matmul_input(input, scale)
input, scale = quantize_int_mm_input(input, scale)
input, weight = check_mats(input, weight)
if bias is not None:
return dequantize_symmetric_with_bias(int_mm_func(input, weight), scale, bias, dtype=return_dtype, result_shape=output_shape)
@@ -57,7 +58,14 @@ def quantized_linear_forward_int8_matmul(self, input: torch.FloatTensor) -> torc
weight = self.weight
scale = self.scale
quantized_weight_shape = self.sdnq_dequantizer.quantized_weight_shape if self.sdnq_dequantizer.is_packed else None
return int8_matmul(input, weight, self.bias, scale, self.svd_up, self.svd_down, quantized_weight_shape, self.sdnq_dequantizer.weights_dtype)
return int8_matmul(
input, weight, scale,
bias=self.bias,
svd_up=self.svd_up,
svd_down=self.svd_down,
quantized_weight_shape=quantized_weight_shape,
weights_dtype=self.sdnq_dequantizer.weights_dtype
)
int8_matmul = compile_func(int8_matmul)
+6 -15
View File
@@ -5,7 +5,6 @@ from diffusers.models.modeling_utils import ModelMixin
from .common import dtype_dict, use_tensorwise_fp8_matmul
from .quantizer import SDNQConfig, sdnq_post_load_quant, prepare_weight_for_matmul, prepare_svd_for_matmul
from .dequantizer import dequantize_symmetric, re_quantize_int8, re_quantize_fp8
from .forward import get_forward_func
from .file_loader import load_files
@@ -177,22 +176,14 @@ def apply_options_to_model(model, dtype: torch.dtype = None, dequantize_fp32: bo
module.svd_down.data = module.svd_down.to(dtype=scale_dtype)
if use_quantized_matmul is not None and use_quantized_matmul != module.sdnq_dequantizer.use_quantized_matmul:
if module.sdnq_dequantizer.weights_dtype in {"int8", "float8_e4m3fn", "float8_e5m2"}:
if use_quantized_matmul and module.sdnq_dequantizer.re_quantize_for_matmul:
scale_dtype = module.scale.dtype
if module.sdnq_dequantizer.weights_dtype == "int8":
module.weight.data, module.scale.data = re_quantize_int8(dequantize_symmetric(module.weight, module.scale, dtype=torch.float32, result_shape=module.sdnq_dequantizer.result_shape))
module.scale.data = module.scale.to(dtype=scale_dtype)
else:
is_e5 = bool(module.sdnq_dequantizer.weights_dtype == "float8_e5m2")
module.weight.data, module.scale.data = re_quantize_fp8(dequantize_symmetric(module.weight, module.scale, dtype=torch.float32, result_shape=module.sdnq_dequantizer.result_shape), is_e5=is_e5)
if use_tensorwise_fp8_matmul:
module.scale.data = module.scale.to(dtype=scale_dtype)
elif not module.sdnq_dequantizer.re_quantize_for_matmul:
module.scale.t_()
module.weight.t_()
if not module.sdnq_dequantizer.re_quantize_for_matmul:
module.scale.t_()
module.weight.t_()
if use_quantized_matmul:
module.weight.data = prepare_weight_for_matmul(module.weight)
else:
module.scale.data = module.scale.contiguous()
module.weight.data = module.weight.contiguous()
if module.svd_up is not None:
module.svd_up.data, module.svd_down.data = prepare_svd_for_matmul(module.svd_up.t_(), module.svd_down.t_(), use_quantized_matmul)
module.sdnq_dequantizer.use_quantized_matmul = use_quantized_matmul
+49 -28
View File
@@ -16,7 +16,7 @@ from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
from modules import devices, shared
from .common import dtype_dict, common_skip_keys, module_skip_keys_dict, accepted_weights, allowed_types, linear_types, conv_types, conv_transpose_types, compile_func, use_tensorwise_fp8_matmul, use_contiguous_mm
from .common import dtype_dict, common_skip_keys, module_skip_keys_dict, accepted_weight_dtypes, accepted_matmul_dtypes, allowed_types, linear_types, conv_types, conv_transpose_types, compile_func, use_tensorwise_fp8_matmul, use_contiguous_mm
from .dequantizer import SDNQDequantizer, dequantize_sdnq_model
from .packed_int import pack_int_symetric, pack_int_asymetric
from .forward import get_forward_func
@@ -189,7 +189,7 @@ def add_module_skip_keys(model, modules_to_not_convert: List[str] = None, module
@devices.inference_context()
def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int8", torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, use_svd=False, use_quantized_matmul=False, use_stochastic_rounding=False, dequantize_fp32=False, param_name=None): # pylint: disable=unused-argument
def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int8", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, use_svd=False, use_quantized_matmul=False, use_stochastic_rounding=False, dequantize_fp32=False, param_name=None): # pylint: disable=unused-argument
num_of_groups = 1
is_conv_type = False
is_conv_transpose_type = False
@@ -200,8 +200,19 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
if torch_dtype is None:
torch_dtype = weight.dtype
if dtype_dict[weights_dtype]["num_bits"] > 8:
use_quantized_matmul = False
if quantized_matmul_dtype is None:
if dtype_dict[weights_dtype]["is_integer"]:
quantized_matmul_dtype = "int8"
elif dtype_dict[weights_dtype]["num_bits"] == 8:
quantized_matmul_dtype = "float8_e4m3fn"
else:
quantized_matmul_dtype = "float16"
re_quantize_for_matmul = bool(
dtype_dict[weights_dtype]["is_unsigned"]
or dtype_dict[weights_dtype]["is_integer"] != dtype_dict[quantized_matmul_dtype]["is_integer"]
or dtype_dict[weights_dtype]["num_bits"] > dtype_dict[quantized_matmul_dtype]["num_bits"]
)
if layer_class_name in conv_types:
if dtype_dict[weights_dtype]["num_bits"] < 4:
@@ -211,9 +222,8 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
output_channel_size, channel_size = weight.shape[:2]
if use_quantized_matmul:
use_quantized_matmul = channel_size >= 32 and output_channel_size >= 32
if use_quantized_matmul and not dtype_dict[weights_dtype]["is_integer"]:
use_quantized_matmul = output_channel_size % 16 == 0 and channel_size % 16 == 0
if use_quantized_matmul and dtype_dict[weights_dtype]["num_bits"] == 8:
use_quantized_matmul = use_quantized_matmul and output_channel_size % 16 == 0 and channel_size % 16 == 0
if use_quantized_matmul and not re_quantize_for_matmul and not dtype_dict[weights_dtype]["is_packed"]:
result_shape = weight.shape
weight = weight.flatten(1,-1)
reduction_axes = -1
@@ -230,14 +240,10 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
try:
output_channel_size, channel_size = weight.shape
except Exception as e:
raise ValueError(f"SDNQ: param_name={param_name} layer_class_name={layer_class_name} weight_shape={weight.shape} weights_dtype={weights_dtype} unsupported") from e
raise ValueError(f"SDNQ: param_name={param_name} layer_class_name={layer_class_name} weight_shape={weight.shape} weights_dtype={weights_dtype} quantized_matmul_dtype={quantized_matmul_dtype} unsupported") from e
if use_quantized_matmul:
use_quantized_matmul = channel_size >= 32 and output_channel_size >= 32
if use_quantized_matmul:
if dtype_dict[weights_dtype]["is_integer"]:
use_quantized_matmul = output_channel_size % 8 == 0 and channel_size % 8 == 0
else:
use_quantized_matmul = output_channel_size % 16 == 0 and channel_size % 16 == 0
use_quantized_matmul = use_quantized_matmul and output_channel_size % 16 == 0 and channel_size % 16 == 0
else:
if weight.ndim > 1:
output_channel_size, channel_size = weight.shape[-2:]
@@ -259,14 +265,12 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
svd_up, svd_down = None, None
if group_size == 0:
if use_quantized_matmul and dtype_dict[weights_dtype]["num_bits"] >= 6:
if use_quantized_matmul and not re_quantize_for_matmul and dtype_dict[weights_dtype]["num_bits"] >= 6:
group_size = -1
elif is_linear_type:
group_size = 2 ** ((2 if svd_up is None else 3) + dtype_dict[weights_dtype]["num_bits"])
else:
group_size = 2 ** ((1 if svd_up is None else 2) + dtype_dict[weights_dtype]["num_bits"])
elif use_quantized_matmul and dtype_dict[weights_dtype]["num_bits"] == 8:
group_size = -1 # override user value, re-quantizing 8bit into 8bit is pointless
elif group_size != -1 and not is_linear_type:
group_size = max(group_size // 2, 1)
@@ -314,8 +318,12 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
weight, scale, zero_point = quantize_weight(weight, reduction_axes, weights_dtype)
if (
not dequantize_fp32
and not (use_quantized_matmul and not dtype_dict[weights_dtype]["is_integer"] and not use_tensorwise_fp8_matmul)
and not (weights_dtype == "uint16" and torch_dtype == torch.float16) # uint16 range is larger than fp16, fp16 will cause NaN on dequant
and dtype_dict[weights_dtype]["num_bits"] <= 8
and not (
use_quantized_matmul
and not dtype_dict[quantized_matmul_dtype]["is_integer"]
and (not use_tensorwise_fp8_matmul or dtype_dict[quantized_matmul_dtype]["num_bits"] == 16)
)
):
scale = scale.to(dtype=torch_dtype)
if zero_point is not None:
@@ -324,7 +332,7 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
svd_up = svd_up.to(dtype=torch_dtype)
svd_down = svd_down.to(dtype=torch_dtype)
re_quantize_for_matmul = (num_of_groups > 1 or zero_point is not None)
re_quantize_for_matmul = re_quantize_for_matmul or num_of_groups > 1
if use_quantized_matmul and not re_quantize_for_matmul:
scale.t_()
weight.t_()
@@ -339,6 +347,7 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
original_stride=original_stride,
quantized_weight_shape=weight.shape,
weights_dtype=weights_dtype,
quantized_matmul_dtype=quantized_matmul_dtype,
group_size=group_size,
svd_rank=svd_rank,
svd_steps=svd_steps,
@@ -356,11 +365,11 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
else:
weight = weight.to(dtype=dtype_dict[weights_dtype]["torch_dtype"])
return weight, scale, zero_point, svd_up, svd_down, sdnq_dequantizer, use_quantized_matmul
return weight, scale, zero_point, svd_up, svd_down, sdnq_dequantizer
@devices.inference_context()
def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, use_svd=False, quant_conv=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, use_stochastic_rounding=False, dequantize_fp32=False, non_blocking=False, quantization_device=None, return_device=None, param_name=None): # pylint: disable=unused-argument
def sdnq_quantize_layer(layer, weights_dtype="int8", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, use_svd=False, quant_conv=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, use_stochastic_rounding=False, dequantize_fp32=False, non_blocking=False, quantization_device=None, return_device=None, param_name=None): # pylint: disable=unused-argument
layer_class_name = layer.__class__.__name__
if layer_class_name in conv_transpose_types or layer_class_name in conv_types:
if not quant_conv:
@@ -379,11 +388,11 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_siz
layer.scale, layer.zero_point,
layer.svd_up, layer.svd_down,
layer.sdnq_dequantizer,
use_quantized_matmul,
) = sdnq_quantize_layer_weight(
layer.weight,
layer_class_name=layer_class_name,
weights_dtype=weights_dtype,
quantized_matmul_dtype=quantized_matmul_dtype,
torch_dtype=torch_dtype,
group_size=group_size,
svd_rank=svd_rank,
@@ -404,13 +413,13 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_siz
layer.svd_down = torch.nn.Parameter(layer.svd_down.to(return_device, non_blocking=non_blocking), requires_grad=False)
layer = layer.to(return_device, non_blocking=non_blocking)
layer.forward = get_forward_func(layer_class_name, use_quantized_matmul, layer.sdnq_dequantizer.is_integer)
layer.forward = get_forward_func(layer_class_name, layer.sdnq_dequantizer.quantized_matmul_dtype, layer.sdnq_dequantizer.use_quantized_matmul)
layer.forward = layer.forward.__get__(layer, layer.__class__)
return layer
@devices.inference_context()
def apply_sdnq_to_module(model, weights_dtype="int8", torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, use_svd=False, quant_conv=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, use_stochastic_rounding=False, dequantize_fp32=False, non_blocking=False, quantization_device=None, return_device=None, modules_to_not_convert: List[str] = None, modules_dtype_dict: Dict[str, List[str]] = None, full_param_name=""): # pylint: disable=unused-argument
def apply_sdnq_to_module(model, weights_dtype="int8", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, use_svd=False, quant_conv=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, use_stochastic_rounding=False, dequantize_fp32=False, non_blocking=False, quantization_device=None, return_device=None, modules_to_not_convert: List[str] = None, modules_dtype_dict: Dict[str, List[str]] = None, full_param_name=""): # pylint: disable=unused-argument
has_children = list(model.children())
if not has_children:
return model
@@ -434,6 +443,7 @@ def apply_sdnq_to_module(model, weights_dtype="int8", torch_dtype=None, group_si
setattr(model, module_name, sdnq_quantize_layer(
module,
weights_dtype=get_minimum_dtype(weights_dtype, param_name, modules_dtype_dict),
quantized_matmul_dtype=quantized_matmul_dtype,
torch_dtype=torch_dtype,
group_size=group_size,
svd_rank=svd_rank,
@@ -452,6 +462,7 @@ def apply_sdnq_to_module(model, weights_dtype="int8", torch_dtype=None, group_si
setattr(model, module_name, apply_sdnq_to_module(
module,
weights_dtype=weights_dtype,
quantized_matmul_dtype=quantized_matmul_dtype,
torch_dtype=torch_dtype,
group_size=group_size,
svd_rank=svd_rank,
@@ -476,6 +487,7 @@ def apply_sdnq_to_module(model, weights_dtype="int8", torch_dtype=None, group_si
def sdnq_post_load_quant(
model: torch.nn.Module,
weights_dtype: str = "int8",
quantized_matmul_dtype: str = None,
torch_dtype: torch.dtype = None,
group_size: int = 0,
svd_rank: int = 32,
@@ -507,6 +519,7 @@ def sdnq_post_load_quant(
model = apply_sdnq_to_module(
model,
weights_dtype=weights_dtype,
quantized_matmul_dtype=quantized_matmul_dtype,
torch_dtype=torch_dtype,
group_size=group_size,
svd_rank=svd_rank,
@@ -663,6 +676,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
layer = sdnq_quantize_layer(
layer,
weights_dtype=weights_dtype,
quantized_matmul_dtype=self.quantization_config.quantized_matmul_dtype,
torch_dtype=torch_dtype,
group_size=self.quantization_config.group_size,
svd_rank=self.quantization_config.svd_rank,
@@ -712,7 +726,6 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
quantization_config_dict.pop("add_skip_keys", None)
quantization_config_dict.pop("use_static_quantization", None)
quantization_config_dict.pop("use_stochastic_rounding", None)
quantization_config_dict.pop("quantized_matmul_dtype", None)
quantization_config_dict.pop("use_grad_ckpt", None)
quantization_config_dict.pop("is_training", None)
with init_empty_weights():
@@ -789,6 +802,10 @@ class SDNQConfig(QuantizationConfigMixin):
weights_dtype (`str`, *optional*, defaults to `"int8"`):
The target dtype for the weights after quantization. Supported values are:
("int16", "int8", "int7", "int6", "int5", "int4", "int3", "int2", "uint16", "uint8", "uint7", "uint6", "uint5", "uint4", "uint3", "uint2", "uint1", "bool", "float16", "float8_e4m3fn", "float8_e4m3fnuz", "float8_e5m2", "float8_e5m2fnuz")
quantized_matmul_dtype (`str`, *optional*, defaults to `None`):
The target dtype for quantized matmul.
`None` will use "int8" with integer weight dtypes and "float8_e4m3fn" or "float16" with float weight dtypes.
Supported values are: ("int8", "float8_e4m3fn", "float16")
group_size (`int`, *optional*, defaults to `0`):
Used to decide how many elements of a tensor will share the same quantization group.
group_size = 0 will automatically select a group size based on weights_dtype.
@@ -804,6 +821,8 @@ class SDNQConfig(QuantizationConfigMixin):
Enabling this option will use quantized INT8 or FP8 MatMul instead of BF16 / FP16.
use_quantized_matmul_conv (`bool`, *optional*, defaults to `False`):
Same as use_quantized_matmul_conv but for the convolutional layers with UNets like SDXL.
use_stochastic_rounding (`bool`, *optional*, defaults to `False`):
Enabling this option will use stochastic rounding on the quantization step.
dequantize_fp32 (`bool`, *optional*, defaults to `False`):
Enabling this option will use FP32 on the dequantization step.
non_blocking (`bool`, *optional*, defaults to `False`):
@@ -824,7 +843,7 @@ class SDNQConfig(QuantizationConfigMixin):
def __init__( # pylint: disable=super-init-not-called
self,
weights_dtype: str = "int8",
quantized_matmul_dtype: str = "int8",
quantized_matmul_dtype: str = None,
group_size: int = 0,
svd_rank: int = 32,
svd_steps: int = 8,
@@ -876,8 +895,10 @@ class SDNQConfig(QuantizationConfigMixin):
r"""
Safety checker that arguments are correct
"""
if self.weights_dtype not in accepted_weights:
raise ValueError(f"SDNQ only support weights in {accepted_weights} but found {self.weights_dtype}")
if self.weights_dtype not in accepted_weight_dtypes:
raise ValueError(f"SDNQ only support weight dtypes in {accepted_weight_dtypes} but found {self.weights_dtype}")
if self.quantized_matmul_dtype is not None and self.quantized_matmul_dtype not in accepted_matmul_dtypes:
raise ValueError(f"SDNQ only support quantized matmul dtypes in {accepted_matmul_dtypes} but found {self.quantized_matmul_dtype}")
if self.modules_to_not_convert is None:
self.modules_to_not_convert = []
+76 -7
View File
@@ -56,13 +56,13 @@ def get_autotune_config():
@triton.autotune(configs=get_autotune_config(), key=['M', 'N', 'K', 'stride_bk'])
@triton.jit
def int_mm_kernel(
a_ptr, b_ptr, c_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, #
GROUP_SIZE_M: tl.constexpr
a_ptr, b_ptr, c_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,
GROUP_SIZE_M: tl.constexpr,
):
pid = tl.program_id(axis=0)
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
@@ -120,3 +120,72 @@ def int_mm(a, b):
c.stride(0), c.stride(1),
)
return c
@triton.autotune(configs=get_autotune_config(), key=['M', 'N', 'K', 'stride_bk'])
@triton.jit
def fp_mm_kernel(
a_ptr, b_ptr, c_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,
GROUP_SIZE_M: tl.constexpr,
):
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
offs_k = tl.arange(0, BLOCK_SIZE_K)
a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak)
b_ptrs = b_ptr + (offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0)
b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k * BLOCK_SIZE_K, other=0.0)
accumulator = tl.dot(a, b, accumulator, out_dtype=tl.float32)
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += BLOCK_SIZE_K * stride_bk
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, accumulator, mask=c_mask)
def fp_mm(a, b):
assert a.shape[1] == b.shape[0], "Incompatible dimensions"
assert a.is_contiguous(), "Matrix A must be contiguous"
M, K = a.shape
K, N = b.shape
c = torch.empty((M, N), device=a.device, dtype=torch.float32)
def grid(META):
return (triton.cdiv(M, META['BLOCK_SIZE_M']) * triton.cdiv(N, META['BLOCK_SIZE_N']), )
fp_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),
)
return c
+4 -1
View File
@@ -63,7 +63,8 @@ restricted_opts = {
}
resize_modes = ["None", "Fixed", "Crop", "Fill", "Outpaint", "Context aware"]
max_workers = 12
sdnq_quant_modes = ["int8", "float8_e4m3fn", "int7", "int6", "int5", "uint4", "uint3", "uint2", "float8_e5m2", "float8_e4m3fnuz", "float8_e5m2fnuz", "uint8", "uint7", "uint6", "uint5", "int4", "int3", "int2", "uint1"]
sdnq_quant_modes = ["int8", "float8_e4m3fn", "int7", "int6", "int5", "uint4", "uint3", "uint2", "float8_e5m2", "float8_e4m3fnuz", "float8_e5m2fnuz", "float16", "int16", "uint16", "uint8", "uint7", "uint6", "uint5", "int4", "int3", "int2", "uint1"]
sdnq_matmul_modes = ["auto", "int8", "float8_e4m3fn", "float16"]
default_hfcache_dir = os.environ.get("SD_HFCACHEDIR", None) or os.path.join(paths.models_path, 'huggingface')
state = shared_state.State()
@@ -201,7 +202,9 @@ options_templates.update(options_section(("quantization", "Model Quantization"),
"sdnq_quantize_weights": OptionInfo([], "Quantization enabled", gr.CheckboxGroup, {"choices": ["Model", "TE", "LLM", "Control", "VAE"]}),
"sdnq_quantize_mode": OptionInfo("auto", "Quantization mode", gr.Dropdown, {"choices": ["auto", "pre", "post"]}),
"sdnq_quantize_weights_mode": OptionInfo("int8", "Quantization type", gr.Dropdown, {"choices": sdnq_quant_modes}),
"sdnq_quantize_matmul_mode": OptionInfo("auto", "Quantized MatMul type", gr.Dropdown, {"choices": sdnq_matmul_modes}),
"sdnq_quantize_weights_mode_te": OptionInfo("Same as model", "Quantization type for Text Encoders", gr.Dropdown, {"choices": ['Same as model'] + sdnq_quant_modes}),
"sdnq_quantize_matmul_mode_te": OptionInfo("Same as model", "Quantized MatMul type for Text Encoders", gr.Dropdown, {"choices": ['Same as model'] + sdnq_matmul_modes}),
"sdnq_modules_to_not_convert": OptionInfo("", "Modules to not convert"),
"sdnq_modules_dtype_dict": OptionInfo("{}", "Modules dtype dict"),
"sdnq_quantize_weights_group_size": OptionInfo(0, "Group size", gr.Slider, {"minimum": -1, "maximum": 4096, "step": 1}),