mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
Improve realquant gemm impl (#368)
Signed-off-by: Chenjie Luo <chenjiel@nvidia.com> Signed-off-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com> Co-authored-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com>
This commit is contained in:
co-authored by
Keval Morabia
parent
576fbdc61e
commit
ad091e884e
@@ -11,7 +11,7 @@
|
||||
:recursive:
|
||||
{% for item in modules %}
|
||||
{% set full_item = fullname + '.' + item.split('.')[-1] %}
|
||||
{% if '.plugins.' not in full_item or full_item == 'modelopt.torch.opt.plugins.huggingface' %}
|
||||
{% if ('.plugins.' not in full_item or full_item == 'modelopt.torch.opt.plugins.huggingface') and full_item != 'modelopt.torch.quantization.backends.fp8_per_tensor_gemm' %}
|
||||
{{ full_item }}
|
||||
{% endif %}
|
||||
{%- endfor %}
|
||||
|
||||
@@ -15,6 +15,5 @@
|
||||
|
||||
"""Quantization backends."""
|
||||
|
||||
from .fp8_per_tensor_gemm import *
|
||||
from .gemm_registry import *
|
||||
from .nvfp4_gemm import *
|
||||
|
||||
@@ -32,35 +32,38 @@ FP8_MIN = torch.finfo(torch.float8_e4m3fn).min
|
||||
FP8_MAX = torch.finfo(torch.float8_e4m3fn).max
|
||||
|
||||
|
||||
@torch.compile(dynamic=True)
|
||||
def _to_fp8(x, scale):
|
||||
return (x / scale).clamp(FP8_MIN, FP8_MAX).to(torch.float8_e4m3fn)
|
||||
|
||||
|
||||
@torch.compile(dynamic=True)
|
||||
def _fp8_gemm_impl(input, weight_fp8, scale_a, scale_b, bias=None):
|
||||
input_shape = input.shape
|
||||
input_fp8 = _to_fp8(input, scale_a).reshape(-1, input_shape[-1])
|
||||
weight_fp8_t = weight_fp8.reshape(-1, weight_fp8.shape[-1]).t()
|
||||
output = torch._scaled_mm(
|
||||
input_fp8,
|
||||
weight_fp8_t,
|
||||
scale_a=scale_a,
|
||||
scale_b=scale_b,
|
||||
bias=bias,
|
||||
out_dtype=input.dtype,
|
||||
use_fast_accum=True,
|
||||
)
|
||||
return output.reshape(*input_shape[:-1], output.shape[-1])
|
||||
|
||||
|
||||
def fp8_per_tensor_gemm(quant_module, input, bias=None):
|
||||
"""GEMM function for fp8 per tensor quantization."""
|
||||
|
||||
@torch.compile(dynamic=True)
|
||||
def _to_fp8(x, scale):
|
||||
return (x / scale).clamp(FP8_MIN, FP8_MAX).to(torch.float8_e4m3fn)
|
||||
|
||||
@torch.compile(dynamic=True)
|
||||
def _fp8_gemm_impl(input, weight_fp8, scale_a, scale_b, bias=None):
|
||||
input_shape = input.shape
|
||||
input_fp8 = _to_fp8(input, scale_a).reshape(-1, input_shape[-1])
|
||||
weight_fp8_t = weight_fp8.reshape(-1, weight_fp8.shape[-1]).t()
|
||||
output = torch._scaled_mm(
|
||||
input_fp8,
|
||||
weight_fp8_t,
|
||||
scale_a=scale_a,
|
||||
scale_b=scale_b,
|
||||
bias=bias,
|
||||
out_dtype=input.dtype,
|
||||
use_fast_accum=True,
|
||||
)
|
||||
return output.reshape(*input_shape[:-1], output.shape[-1])
|
||||
|
||||
cached_scale_a = (
|
||||
hasattr(quant_module, "_scale_a") and quant_module.input_quantizer.amax is not None
|
||||
)
|
||||
|
||||
if not cached_scale_a:
|
||||
input_amax = quant_module.input_quantizer.amax or reduce_amax(input)
|
||||
input_amax = quant_module.input_quantizer.amax
|
||||
if input_amax is None:
|
||||
input_amax = reduce_amax(input)
|
||||
assert input_amax != 0
|
||||
quant_module._scale_a = (input_amax.float() / 448.0).to(device=input.device)
|
||||
|
||||
@@ -69,7 +72,9 @@ def fp8_per_tensor_gemm(quant_module, input, bias=None):
|
||||
)
|
||||
|
||||
if not cached_scale_b:
|
||||
weight_amax = quant_module.weight_quantizer.amax or reduce_amax(quant_module.weight)
|
||||
weight_amax = quant_module.weight_quantizer.amax
|
||||
if weight_amax is None:
|
||||
weight_amax = reduce_amax(quant_module.weight)
|
||||
assert weight_amax != 0
|
||||
quant_module._scale_b = (weight_amax.float() / 448.0).to(device=quant_module.weight.device)
|
||||
|
||||
@@ -146,9 +151,9 @@ class Fp8PerTensorLinear(Function):
|
||||
ctx.save_for_backward(
|
||||
input_tensor if weight.requires_grad else None,
|
||||
weight if input_tensor.requires_grad else None,
|
||||
torch.empty(0, dtype=torch.uint8) if bias is not None and bias.requires_grad else None,
|
||||
getattr(quant_module.weight_quantizer, "_scale", None),
|
||||
)
|
||||
ctx.compute_bias_grad = bias is not None and bias.requires_grad
|
||||
ctx.block_sizes = getattr(quant_module.weight_quantizer, "_block_sizes", None)
|
||||
|
||||
ctx.allreduce_dgrad = allreduce_dgrad
|
||||
@@ -166,7 +171,7 @@ class Fp8PerTensorLinear(Function):
|
||||
dequantize it to compute the input gradient. If the weight is not compressed, we will save
|
||||
the unquantized weight and use it directly to compute the input gradient.
|
||||
"""
|
||||
input_tensor, weight, compute_bias_grad, scale = ctx.saved_tensors
|
||||
input_tensor, weight, scale = ctx.saved_tensors
|
||||
grad_input = grad_weight = grad_bias = None
|
||||
if weight is not None:
|
||||
if isinstance(weight, QTensorWrapper):
|
||||
@@ -175,8 +180,10 @@ class Fp8PerTensorLinear(Function):
|
||||
weight = weight.dequantize(scale=scale, block_sizes=ctx.block_sizes)
|
||||
grad_input = grad_outputs @ weight
|
||||
if input_tensor is not None:
|
||||
grad_weight = grad_outputs.transpose(-2, 1) @ input_tensor
|
||||
if compute_bias_grad is not None:
|
||||
grad_weight = grad_outputs.reshape(-1, grad_outputs.shape[-1]).T @ input_tensor.reshape(
|
||||
-1, input_tensor.shape[-1]
|
||||
)
|
||||
if ctx.compute_bias_grad:
|
||||
# Sum all dimensions except the last one
|
||||
grad_bias = grad_outputs.sum(dim=list(range(grad_outputs.dim() - 1)))
|
||||
|
||||
|
||||
@@ -139,11 +139,11 @@ class Nvfp4Linear(Function):
|
||||
ctx.save_for_backward(
|
||||
input_tensor if weight.requires_grad else None,
|
||||
weight if input_tensor.requires_grad else None,
|
||||
torch.empty(0, dtype=torch.uint8) if bias is not None and bias.requires_grad else None,
|
||||
getattr(quant_module.weight_quantizer, "_scale", None),
|
||||
getattr(quant_module.weight_quantizer, "_double_scale", None),
|
||||
)
|
||||
|
||||
ctx.compute_bias_grad = bias is not None and bias.requires_grad
|
||||
ctx.allreduce_dgrad = allreduce_dgrad
|
||||
ctx.tp_group = tp_group
|
||||
ret = nvfp4_gemm(quant_module, input_tensor, bias)
|
||||
@@ -158,7 +158,7 @@ class Nvfp4Linear(Function):
|
||||
dequantize it to compute the input gradient. If the weight is not compressed, we will save
|
||||
the unquantized weight and use it directly to compute the input gradient.
|
||||
"""
|
||||
input_tensor, weight, compute_bias_grad, scale, double_scale = ctx.saved_tensors
|
||||
input_tensor, weight, scale, double_scale = ctx.saved_tensors
|
||||
grad_input = grad_weight = grad_bias = None
|
||||
if weight is not None:
|
||||
if isinstance(weight, QTensorWrapper):
|
||||
@@ -173,8 +173,10 @@ class Nvfp4Linear(Function):
|
||||
)
|
||||
grad_input = grad_outputs @ weight
|
||||
if input_tensor is not None:
|
||||
grad_weight = grad_outputs.transpose(-2, -1) @ input_tensor
|
||||
if compute_bias_grad is not None:
|
||||
grad_weight = grad_outputs.reshape(-1, grad_outputs.shape[-1]).T @ input_tensor.reshape(
|
||||
-1, input_tensor.shape[-1]
|
||||
)
|
||||
if ctx.compute_bias_grad is not None:
|
||||
# Sum all dimensions except the last one
|
||||
grad_bias = grad_outputs.sum(dim=list(range(grad_outputs.dim() - 1)))
|
||||
|
||||
|
||||
@@ -148,7 +148,7 @@ class RealQuantLinear(QuantModule):
|
||||
and self.allow_real_quant_gemm
|
||||
)
|
||||
|
||||
def get_real_quant_gemm_impl(self, input, *args, **kwargs) -> bool:
|
||||
def has_real_quant_gemm_impl(self, input, *args, **kwargs) -> bool:
|
||||
"""Get the real quant GEMM implementation base on input arguments."""
|
||||
if not hasattr(self, "_real_quant_gemm_impl"):
|
||||
self._real_quant_gemm_impl = backends.gemm_registry.find_match(
|
||||
@@ -166,20 +166,19 @@ class RealQuantLinear(QuantModule):
|
||||
return super().forward(input, *args, **kwargs)
|
||||
|
||||
# Check if real-quant GEMM is available
|
||||
if self._should_run_real_quant_gemm and input.numel() > 1:
|
||||
# If the input is not quantized, we use the default GEMM.
|
||||
self.get_real_quant_gemm_impl(input, *args, **kwargs)
|
||||
|
||||
if (
|
||||
self._should_run_real_quant_gemm
|
||||
and input.numel() > 1
|
||||
and self.has_real_quant_gemm_impl(input, *args, **kwargs)
|
||||
):
|
||||
# Note: We cache the real-quant GEMM function to avoid matching overhead.
|
||||
# This assumes that the function will not change after the first call.
|
||||
if self._real_quant_gemm_impl:
|
||||
with torch.cuda.nvtx.range("RealQuantLinear gemm"):
|
||||
output = self._real_quant_gemm_impl(
|
||||
self, input, self.weight, self.bias, *args, **kwargs
|
||||
)
|
||||
return (
|
||||
self.output_quantizer(output) if hasattr(self, "output_quantizer") else output
|
||||
assert self._real_quant_gemm_impl is not None
|
||||
with torch.cuda.nvtx.range("RealQuantLinear gemm"):
|
||||
output = self._real_quant_gemm_impl(
|
||||
self, input, self.weight, self.bias, *args, **kwargs
|
||||
)
|
||||
return self.output_quantizer(output) if hasattr(self, "output_quantizer") else output
|
||||
|
||||
# Otherwise, fallback to the default GEMM
|
||||
return super().forward(input, *args, **kwargs)
|
||||
|
||||
@@ -115,7 +115,7 @@ def real_quant_module_set_extra_state(self, state: Any):
|
||||
"""
|
||||
q_tensor_state = state.get("modelopt_q_tensor_state", None)
|
||||
|
||||
if q_tensor_state is not None:
|
||||
if q_tensor_state:
|
||||
q_tensor_metadata = q_tensor_state["metadata"]
|
||||
q_tensor_metadata["shape"] = self.weight.shape
|
||||
q_tensor_data_dtype = q_tensor_state["quantized_data.dtype"]
|
||||
@@ -418,8 +418,10 @@ class _RealQuantMegatronParallelLinear(RealQuantLinear):
|
||||
while the original forward only takes 1 positional argument. We must above the fallback path
|
||||
in RealQuantLinear.forward().
|
||||
"""
|
||||
if self._should_run_real_quant_gemm and self.get_real_quant_gemm_impl(
|
||||
input, *args, **kwargs
|
||||
if (
|
||||
self._should_run_real_quant_gemm
|
||||
and input.numel() > 1
|
||||
and self.has_real_quant_gemm_impl(input, *args, **kwargs)
|
||||
):
|
||||
allreduce_dgrad = kwargs.get("allreduce_dgrad", False)
|
||||
tp_group = kwargs.get("tp_group")
|
||||
|
||||
Reference in New Issue
Block a user