mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 01:04:32 +02:00
SDNQ improve svd and low bit matmul perf
This commit is contained in:
@@ -14,7 +14,11 @@ def dequantize_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, ze
|
||||
result = result.view(result_shape)
|
||||
if svd_up is not None:
|
||||
if skip_quantized_matmul:
|
||||
svd_up, svd_down = svd_up.t(), svd_down.t()
|
||||
svd_up = svd_up.t().contiguous()
|
||||
if use_contiguous_mm:
|
||||
svd_down = svd_down.t().contiguous()
|
||||
else:
|
||||
svd_down = svd_down.contiguous().t()
|
||||
if result.ndim > 2 and weight.ndim > 2: # convs
|
||||
result = result.add_(torch.mm(svd_up, svd_down).unflatten(-1, (*result.shape[1:],)))
|
||||
else:
|
||||
@@ -32,7 +36,11 @@ def dequantize_symmetric(weight: torch.CharTensor, scale: torch.FloatTensor, dty
|
||||
result = result.view(result_shape)
|
||||
if svd_up is not None:
|
||||
if skip_quantized_matmul:
|
||||
svd_up, svd_down = svd_up.t(), svd_down.t()
|
||||
svd_up = svd_up.t().contiguous()
|
||||
if use_contiguous_mm:
|
||||
svd_down = svd_down.t().contiguous()
|
||||
else:
|
||||
svd_down = svd_down.contiguous().t()
|
||||
if result.ndim > 2 and weight.ndim > 2: # convs
|
||||
result = result.add_(torch.mm(svd_up, svd_down).unflatten(-1, (*result.shape[1:],)))
|
||||
else:
|
||||
@@ -71,17 +79,20 @@ def quantize_fp8(input: torch.FloatTensor, dim: int = -1, is_e5: bool = False) -
|
||||
def re_quantize_int8(weight: torch.FloatTensor) -> Tuple[torch.CharTensor, torch.FloatTensor]:
|
||||
if weight.ndim > 2: # convs
|
||||
weight = weight.flatten(1,-1)
|
||||
weight = weight.t()
|
||||
if use_contiguous_mm:
|
||||
weight = weight.contiguous()
|
||||
weight, scale = quantize_int8(weight, dim=0)
|
||||
weight, scale = quantize_int8(weight.t(), dim=-0)
|
||||
weight, scale = weight.contiguous(), scale.contiguous()
|
||||
else:
|
||||
weight, scale = quantize_int8(weight.contiguous(), dim=-1)
|
||||
weight, scale = weight.t_(), scale.t_()
|
||||
return weight, scale
|
||||
|
||||
|
||||
def re_quantize_fp8(weight: torch.FloatTensor, is_e5: bool = False) -> Tuple[torch.CharTensor, torch.FloatTensor]:
|
||||
if weight.ndim > 2: # convs
|
||||
weight = weight.flatten(1,-1)
|
||||
weight, scale = quantize_fp8(weight.t(), dim=0, is_e5=is_e5)
|
||||
weight, scale = quantize_int8(weight.contiguous(), dim=-1, is_e5=is_e5)
|
||||
weight, scale = weight.t_(), scale.t_()
|
||||
if not use_tensorwise_fp8_matmul:
|
||||
scale = scale.to(dtype=torch.float32)
|
||||
return weight, scale
|
||||
@@ -95,11 +106,11 @@ def re_quantize_matmul_symmetric(weight: torch.CharTensor, scale: torch.FloatTen
|
||||
return re_quantize_int8(dequantize_symmetric(weight, scale, scale.dtype, result_shape, svd_up=svd_up, svd_down=svd_down))
|
||||
|
||||
|
||||
def re_quantize_matmul_packed_int_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, result_shape: torch.Size, weights_dtype: str, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None) -> torch.FloatTensor:
|
||||
def re_quantize_matmul_packed_int_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, result_shape: torch.Size, weights_dtype: str, 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, result_shape, svd_up=svd_up, svd_down=svd_down)
|
||||
|
||||
|
||||
def re_quantize_matmul_packed_int_symmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, shape: torch.Size, result_shape: torch.Size, weights_dtype: str, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None) -> torch.FloatTensor:
|
||||
def re_quantize_matmul_packed_int_symmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, shape: torch.Size, result_shape: torch.Size, weights_dtype: str, 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, result_shape, svd_up=svd_up, svd_down=svd_down)
|
||||
|
||||
|
||||
|
||||
@@ -187,6 +187,21 @@ def apply_options_to_model(model, dtype: torch.dtype = None, dequantize_fp32: bo
|
||||
if module.svd_up is not None:
|
||||
module.svd_up.data = module.svd_up.t_()
|
||||
module.svd_down.data = module.svd_down.t_()
|
||||
if use_quantized_matmul:
|
||||
if use_contiguous_mm:
|
||||
module.svd_up.data = module.svd_up.contiguous()
|
||||
module.svd_down.data = module.svd_down.contiguous()
|
||||
else:
|
||||
if svd_up.is_contiguous():
|
||||
module.svd_up.data = module.svd_up.t_().contiguous().t_()
|
||||
if svd_up.is_contiguous():
|
||||
module.svd_down.data = module.svd_down.t_().contiguous().t_()
|
||||
else:
|
||||
svd_up = svd_up.contiguous()
|
||||
if use_contiguous_mm:
|
||||
svd_down = svd_down.contiguous()
|
||||
elif svd_down.is_contiguous():
|
||||
svd_down = svd_down.t_().contiguous().t_()
|
||||
module.sdnq_dequantizer.use_quantized_matmul = use_quantized_matmul
|
||||
module = apply_options_to_model(module, dtype=dtype, dequantize_fp32=dequantize_fp32, use_quantized_matmul=use_quantized_matmul)
|
||||
return model
|
||||
|
||||
@@ -204,6 +204,20 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_siz
|
||||
if use_quantized_matmul:
|
||||
svd_up = svd_up.t_()
|
||||
svd_down = svd_down.t_()
|
||||
if use_contiguous_mm:
|
||||
svd_up = svd_up.contiguous()
|
||||
svd_down = svd_down.contiguous()
|
||||
else:
|
||||
if svd_up.is_contiguous():
|
||||
svd_up = svd_up.t_().contiguous().t_()
|
||||
if svd_down.is_contiguous():
|
||||
svd_down = svd_down.t_().contiguous().t_()
|
||||
else:
|
||||
svd_up = svd_up.contiguous()
|
||||
if use_contiguous_mm:
|
||||
svd_down = svd_down.contiguous()
|
||||
elif svd_down.is_contiguous():
|
||||
svd_down = svd_down.t_().contiguous().t_()
|
||||
except Exception:
|
||||
svd_up, svd_down = None, None
|
||||
else:
|
||||
|
||||
Reference in New Issue
Block a user