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:
Chenjie Luo
2025-09-26 09:37:51 +05:30
committed by GitHub
co-authored by Keval Morabia
parent 576fbdc61e
commit ad091e884e
6 changed files with 57 additions and 48 deletions
@@ -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")