mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[6508436] Fix BF16 FP8 ONNX export (#2314)
### What does this PR do?
Type of change: Bug fix
Fix FP8 ONNX export for BF16 models during real-weight compression
without changing the public API or the `weights_dtype="fp32"` default.
The FP8 exporter preserves BF16 initializer bits when bridging
GraphSurgeon NumPy arrays to Torch, widens BF16 values exactly to FP32
for normalization, and leaves existing FP16/FP32 handling unchanged.
Conv scales and dequantized outputs retain the source dtype, and scales
round upward when needed so serialized values cannot cause FP8 overflow.
`weights_dtype="bf16"` is accepted as a no-op only for FP8-only models
whose floating parameters are all BF16. Registered buffers do not affect
this weight-focused decision and may preserve higher-precision regions
in the exported graph. Unsupported BF16 FP8-to-FP16 and FP32 or
mixed-parameter-to-BF16 conversions are rejected with `ValueError`
before temporary export paths are created. A narrow GraphSurgeon fix
preserves integer BF16 value-info dtypes.
### Usage
```python
onnx_bytes, metadata = get_onnx_bytes_and_metadata(
quantized_fp8_model,
(sample_input,),
weights_dtype="bf16",
onnx_opset=23,
)
```
### Testing
- Seven focused CPU regressions passed: BF16 QDQ compression and integer
dtype handling, BF16-to-BF16 and FP32-to-FP16 Conv/Linear export, and
four unsupported-conversion cases.
- QDQ utilities: 31 passed; pytest 2.25s, wall 29.88s.
- FP8 MHA exporter: 6 passed; pytest 2.05s, wall 32.71s.
- Torch deploy utilities: 51 passed; pytest 8.94s, wall 25.74s.
- Torch ONNX CPU export: 36 passed; pytest 4.85s, wall 32.38s.
- Changed-file pre-commit hooks: all passed; wall 8.07s.
- Exact-head FP8 BF16 GPU workflow at `f21d62a`: exit code 0; ONNX
checker passed; 6 FP8 initializers, 3 native `DequantizeLinear` nodes,
and 12 BF16 initializers.
- Refreshed GitHub CI at `f21d62a`: 50 passed and 1 skipped. Unit, GPU,
and regression required aggregates and Codecov passed. Two ONNX example
leaves failed because the runner could not load a cuDNN sublibrary;
their dependent example aggregate consequently failed.
### Before your PR is "*Ready for review*"
Make sure you read and follow [Contributor
guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)
and your commits are signed (`git commit -s -S`).
Make sure you read and follow the [Security Best
Practices](https://github.com/NVIDIA/Model-Optimizer/blob/main/SECURITY.md#security-coding-practices-for-contributors)
(e.g. avoiding hardcoded `trust_remote_code=True`, `torch.load(...,
weights_only=False)`, `pickle`, etc.).
- Is this change backward compatible?: ✅
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: N/A
- Did you write any new necessary tests?: ✅
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅
### Additional Information
- TODO: Deliver authoritative `native`/FP32/FP16/BF16 ONNX export across
all quantized formats in follow-up pull requests.
> 🤖 _Generated by Codex (AI agent)._
---------
Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com>
Co-authored-by: Codex <noreply@openai.com>
Co-authored-by: Codex <codex@openai.com>
This commit is contained in:
co-authored by
Codex
Codex
parent
d19925e446
commit
7f7c46d820
@@ -90,6 +90,7 @@ Changelog
|
||||
- Fix ONNX AutoCast failing on models with external initializers larger than 2 GiB.
|
||||
- Avoid querying CUDA/Blackwell capability when ``NVFP4QTensor.quantize`` uses its CPU path or has the optional TensorRT-LLM fast path disabled.
|
||||
- Fix NVFP4 ONNX export to quantize FP4 weights with the published FP8 block scales, matching eager ModelOpt packed weights. Block scales below ``2**-9`` are now clamped to that minimum, and non-finite or negative scales raise an error.
|
||||
- Fix FP8 ONNX export of BF16 models during real-weight compression.
|
||||
- Fix Megatron-Bridge Quantization Aware Distillation of a vision-language model silently discarding the ModelOpt state, so the distilled checkpoint restored no quantizers and exported as an unquantized model. Re-run QAD to regenerate any affected checkpoint.
|
||||
- Fix Megatron-Core HuggingFace export silently omitting fused (grouped GEMM) MoE experts for architectures without an ``experts.linear_fc1`` rule (e.g. ``Qwen3MoeForCausalLM``), which produced a valid-looking checkpoint containing no expert weights. The exporter now raises instead of writing that checkpoint; the scripts also avoid the situation by selecting ``SequentialMLP`` for those architectures.
|
||||
- Fix GatedDeltaNet (Qwen3.5) quantizer exclusions on Megatron-Core: the recipe patterns name the HuggingFace ``linear_attn`` module, so the ``conv1d`` was calibrated and the alpha / beta gate projections were exported in FP8. ``conv1d`` now has a ``self_attention`` alias in the default disabled-quantizer units, and the alpha / beta projections are exported in BF16 (they share Megatron's fused ``in_proj`` quantizer and cannot be disabled by name).
|
||||
|
||||
@@ -17,6 +17,7 @@
|
||||
|
||||
import time
|
||||
|
||||
import ml_dtypes
|
||||
import numpy as np
|
||||
import onnx
|
||||
import onnx_graphsurgeon as gs
|
||||
@@ -33,6 +34,13 @@ _FP8_E4M3_MAX = 448.0
|
||||
_FP8_E4M3_SOFTMAX_SCALE = 1.0 / _FP8_E4M3_MAX
|
||||
|
||||
|
||||
def _torch_from_numpy_for_fp8(array: np.ndarray) -> torch.Tensor:
|
||||
"""Convert a NumPy array to the PyTorch dtype used for FP8 normalization."""
|
||||
if array.dtype == ml_dtypes.bfloat16:
|
||||
return torch.from_numpy(array.view(np.int16)).view(torch.bfloat16).float()
|
||||
return torch.from_numpy(array)
|
||||
|
||||
|
||||
class FP8QuantExporter(ONNXQuantExporter):
|
||||
"""Exporter for FP8 quantization."""
|
||||
|
||||
@@ -48,7 +56,7 @@ class FP8QuantExporter(ONNXQuantExporter):
|
||||
|
||||
@staticmethod
|
||||
def compress_weights(onnx_model: onnx.ModelProto) -> onnx.ModelProto:
|
||||
"""Compresses FP32/FP16 weights to FP8 by folding QDQ nodes to DQ only.
|
||||
"""Compresses FP32/FP16/BF16 weights to FP8 by folding QDQ nodes to DQ only.
|
||||
|
||||
Even though modelopt supports FP8 onnx export, the weights are represented in fp32 + QDQ.
|
||||
The storage is therefore very bad. In this function,
|
||||
@@ -56,7 +64,7 @@ class FP8QuantExporter(ONNXQuantExporter):
|
||||
weights in the output model. TRT custom ops are converted to native ONNX DequantizeLinear.
|
||||
|
||||
Parameters:
|
||||
onnx_model: ONNX model with FP32/FP16 weights and TRT_FP8 QDQ nodes.
|
||||
onnx_model: ONNX model with FP32/FP16/BF16 weights and TRT_FP8 QDQ nodes.
|
||||
|
||||
Returns:
|
||||
ONNX model with FP8 weights and native ONNX DQ nodes for weights (QDQ preserved for activations).
|
||||
@@ -78,8 +86,8 @@ class FP8QuantExporter(ONNXQuantExporter):
|
||||
|
||||
weights = node.inputs[0]
|
||||
scale = node.inputs[1]
|
||||
torch_weights = torch.from_numpy(weights.values)
|
||||
torch_scale = torch.from_numpy(scale.values)
|
||||
torch_weights = _torch_from_numpy_for_fp8(weights.values)
|
||||
torch_scale = _torch_from_numpy_for_fp8(scale.values)
|
||||
quantizer_name = scale.name.rsplit("/", 1)[0]
|
||||
dq_op = node.outputs[0].outputs[0]
|
||||
if dq_op.op != "TRT_FP8DequantizeLinear":
|
||||
@@ -194,14 +202,25 @@ class FP8QuantExporter(ONNXQuantExporter):
|
||||
if any(out.op == "DequantizeLinear" for out in weight_input.outputs):
|
||||
continue
|
||||
|
||||
torch_weights = torch.from_numpy(weight_input.values.copy())
|
||||
torch_weights = _torch_from_numpy_for_fp8(weight_input.values.copy())
|
||||
amax = torch_weights.abs().max().float()
|
||||
if amax == 0:
|
||||
continue
|
||||
scale_val = (amax / _FP8_E4M3_MAX).item()
|
||||
scale_data = np.array(scale_val, dtype=weight_input.values.dtype)
|
||||
# Round up so normalizing by the serialized scale stays within the FP8 range.
|
||||
if scale_data < scale_val:
|
||||
np.nextafter(
|
||||
scale_data,
|
||||
np.array(np.inf, dtype=scale_data.dtype),
|
||||
out=scale_data,
|
||||
)
|
||||
torch_scale = _torch_from_numpy_for_fp8(scale_data)
|
||||
|
||||
# Quantize weights to FP8 (WAR: numpy doesn't support fp8)
|
||||
fp8_data = (torch_weights / scale_val).to(torch.float8_e4m3fn).view(torch.uint8).numpy()
|
||||
fp8_data = (
|
||||
(torch_weights / torch_scale).to(torch.float8_e4m3fn).view(torch.uint8).numpy()
|
||||
)
|
||||
fp8_tensor = onnx.TensorProto()
|
||||
fp8_tensor.data_type = onnx.TensorProto.FLOAT8E4M3FN
|
||||
fp8_tensor.dims.extend(fp8_data.shape)
|
||||
@@ -210,13 +229,16 @@ class FP8QuantExporter(ONNXQuantExporter):
|
||||
node.name + "/weight_quantizer/fp8_weights", LazyValues(fp8_tensor)
|
||||
)
|
||||
|
||||
# Scale in FP16 — DQ output type matches scale dtype, must match activation type
|
||||
scale_constant = gs.Constant(
|
||||
node.name + "/weight_quantizer/scale",
|
||||
np.array(scale_val, dtype=np.float16),
|
||||
scale_data,
|
||||
)
|
||||
|
||||
dq_output = gs.Variable(node.name + "/weight_quantizer/dq_output")
|
||||
dq_output = gs.Variable(
|
||||
node.name + "/weight_quantizer/dq_output",
|
||||
scale_data.dtype,
|
||||
weight_input.values.shape,
|
||||
)
|
||||
dq_node = gs.Node(
|
||||
op="DequantizeLinear",
|
||||
name=node.name + "/weight_quantizer/DequantizeLinear",
|
||||
|
||||
@@ -101,9 +101,15 @@ def _export_value_info_proto(tensor: gs.Variable, do_type_check: bool) -> onnx.V
|
||||
)
|
||||
|
||||
if tensor.dtype is not None:
|
||||
dtype = getattr(
|
||||
tensor, "explicit_dtype", onnx.helper.np_dtype_to_tensor_dtype(np.dtype(tensor.dtype))
|
||||
)
|
||||
dtype = getattr(tensor, "explicit_dtype", None)
|
||||
if dtype is None:
|
||||
dtype = tensor.dtype
|
||||
if isinstance(dtype, (int, np.integer)):
|
||||
dtype = int(dtype)
|
||||
if dtype not in onnx.TensorProto.DataType.values():
|
||||
raise ValueError(f"Unknown ONNX tensor dtype for {tensor.name}: {dtype}")
|
||||
else:
|
||||
dtype = onnx.helper.np_dtype_to_tensor_dtype(np.dtype(dtype))
|
||||
onnx_tensor = onnx.helper.make_tensor_value_info(tensor.name, dtype, tensor.shape)
|
||||
else:
|
||||
onnx_tensor = onnx.helper.make_empty_tensor_value_info(tensor.name)
|
||||
|
||||
@@ -507,14 +507,17 @@ def get_onnx_bytes_and_metadata(
|
||||
`torch.onnx.export <https://pytorch.org/docs/stable/onnx.html#torch.onnx.export>`_.
|
||||
onnx_opset: The onnx opset version to use for exporting the model.
|
||||
dq_only: If True, the exported onnx model is converted to a dq_only model.
|
||||
weights_dtype: The dtype of the weights in the onnx model.
|
||||
weights_dtype: Requested high-precision dtype for exported weights. For an FP8 model,
|
||||
``"bf16"`` is accepted only when every floating parameter is already BF16. This is
|
||||
a weight-focused no-op, not a graph-wide conversion: floating buffers are not
|
||||
considered for eligibility and may preserve higher-precision regions.
|
||||
|
||||
Returns:
|
||||
bytes: Onnx model in bytes.
|
||||
ModelMetadata: The model's meta data.
|
||||
|
||||
Raises:
|
||||
ValueError: If nn.Module is not passed as model.
|
||||
ValueError: If model is not an nn.Module or the requested precision conversion is unsupported.
|
||||
"""
|
||||
if not isinstance(model, nn.Module):
|
||||
raise ValueError("Only PyTorch model compilation is supported.")
|
||||
@@ -527,6 +530,22 @@ def get_onnx_bytes_and_metadata(
|
||||
if isinstance(model, (DataParallel, DistributedDataParallel)):
|
||||
model = model.module
|
||||
|
||||
source_parameter_dtypes = {
|
||||
parameter.dtype for parameter in model.parameters() if parameter.is_floating_point()
|
||||
}
|
||||
source_parameter_dtype_names = ", ".join(sorted(map(str, source_parameter_dtypes))) or "none"
|
||||
uses_fp4 = is_fp4_quantized(model)
|
||||
uses_mxfp8 = is_mxfp8_quantized(model)
|
||||
uses_fp8 = is_fp8_quantized(model)
|
||||
uses_int8 = is_int8_quantized(model)
|
||||
uses_other_unsupported_quantizer = is_int4_quantized(model) or uses_mxfp8 or uses_int8
|
||||
is_bf16_fp8_noop = (
|
||||
weights_dtype == "bf16"
|
||||
and source_parameter_dtypes == {torch.bfloat16}
|
||||
and uses_fp8
|
||||
and not (uses_fp4 or uses_other_unsupported_quantizer)
|
||||
)
|
||||
|
||||
# Standardize model args and also tensorize them so they also appear in the onnx graph!
|
||||
# Floats/ints are tensorized when they are provided, but not tensorized when they are not
|
||||
# provided which is somewhat inconsistent (we always tensorize them!)
|
||||
@@ -549,11 +568,7 @@ def get_onnx_bytes_and_metadata(
|
||||
input_none_names = list(set(tree_spec_input.names) - set(input_names))
|
||||
|
||||
use_torch_autocast = not (
|
||||
is_fp4_quantized(model)
|
||||
or is_mxfp8_quantized(model)
|
||||
or is_fp8_quantized(model)
|
||||
or is_int8_quantized(model)
|
||||
or weights_dtype == "fp32"
|
||||
uses_fp4 or uses_mxfp8 or uses_fp8 or uses_int8 or weights_dtype == "fp32"
|
||||
)
|
||||
autocast = torch.autocast("cuda") if use_torch_autocast else nullcontext()
|
||||
|
||||
@@ -575,6 +590,22 @@ def get_onnx_bytes_and_metadata(
|
||||
)
|
||||
return onnx_model.to_bytes(), model_metadata
|
||||
|
||||
if weights_dtype == "fp16" and uses_fp8 and torch.bfloat16 in source_parameter_dtypes:
|
||||
raise ValueError(
|
||||
"Converting a BF16 FP8 ONNX graph to FP16 is not supported yet "
|
||||
f"(source parameter dtypes: {source_parameter_dtype_names})"
|
||||
)
|
||||
|
||||
if (
|
||||
weights_dtype == "bf16"
|
||||
and (uses_fp8 or uses_other_unsupported_quantizer)
|
||||
and not is_bf16_fp8_noop
|
||||
):
|
||||
raise ValueError(
|
||||
"Converting a quantized ONNX graph to BF16 is not supported yet "
|
||||
f"(source parameter dtypes: {source_parameter_dtype_names})"
|
||||
)
|
||||
|
||||
# Export onnx model from pytorch model
|
||||
# As the maximum size of protobuf is 2GB, we cannot use io.BytesIO() buffer during export.
|
||||
model_name = model_name or model.__class__.__name__
|
||||
@@ -583,16 +614,12 @@ def get_onnx_bytes_and_metadata(
|
||||
|
||||
# Configure quantizers if the model is quantized in NVFP4 or MXFP8 mode
|
||||
quantizer_context = (
|
||||
configure_linear_module_onnx_quantizers(model)
|
||||
if is_fp4_quantized(model) or is_mxfp8_quantized(model)
|
||||
else nullcontext()
|
||||
configure_linear_module_onnx_quantizers(model) if uses_fp4 or uses_mxfp8 else nullcontext()
|
||||
)
|
||||
# Disable FP8 Conv weight quantizers: TorchScript custom ops produce outputs with
|
||||
# unknown shapes, causing _convolution symbolic to fail. Conv weights are quantized
|
||||
# to FP8 in post-processing by FP8QuantExporter instead.
|
||||
conv_wq_context = (
|
||||
_disable_fp8_conv_weight_quantizers(model) if is_fp8_quantized(model) else nullcontext()
|
||||
)
|
||||
conv_wq_context = _disable_fp8_conv_weight_quantizers(model) if uses_fp8 else nullcontext()
|
||||
with torch.inference_mode(), autocast, quantizer_context, conv_wq_context:
|
||||
additional_kwargs = {}
|
||||
if not dynamo_export:
|
||||
@@ -634,14 +661,8 @@ def get_onnx_bytes_and_metadata(
|
||||
if dq_only:
|
||||
onnx_opt_graph = qdq_to_dq(onnx_opt_graph)
|
||||
|
||||
if weights_dtype in ["fp16", "bf16"]:
|
||||
if (
|
||||
is_int4_quantized(model)
|
||||
or is_mxfp8_quantized(model)
|
||||
or is_fp8_quantized(model)
|
||||
or is_int8_quantized(model)
|
||||
):
|
||||
assert weights_dtype == "fp16", "BF16 + MXFP8/INT4 mixed precision is not supported yet"
|
||||
if weights_dtype in ["fp16", "bf16"] and not is_bf16_fp8_noop:
|
||||
if uses_other_unsupported_quantizer or uses_fp8:
|
||||
onnx_opt_graph = convert_float_to_float16(
|
||||
onnx_opt_graph,
|
||||
keep_io_types=False,
|
||||
@@ -665,7 +686,7 @@ def get_onnx_bytes_and_metadata(
|
||||
onnx_opt_graph = remove_redundant_casts(onnx_opt_graph)
|
||||
|
||||
# Remove Cast nodes around Q/DQ for optimal TRT fusion
|
||||
if is_fp8_quantized(model):
|
||||
if uses_fp8:
|
||||
onnx_opt_graph = fold_q_fp16_to_fp32_casts(onnx_opt_graph)
|
||||
onnx_opt_graph = fold_dq_fp32_to_fp16_casts(onnx_opt_graph)
|
||||
|
||||
|
||||
@@ -15,14 +15,22 @@
|
||||
|
||||
import warnings
|
||||
|
||||
import ml_dtypes
|
||||
import numpy as np
|
||||
import onnx
|
||||
import onnx_graphsurgeon as gs
|
||||
import onnxruntime as ort
|
||||
import pytest
|
||||
from onnx import TensorProto, helper, numpy_helper
|
||||
|
||||
from modelopt.onnx.export import INT4QuantExporter, MXFP8QuantExporter, NVFP4QuantExporter
|
||||
from modelopt.onnx.export import (
|
||||
FP8QuantExporter,
|
||||
INT4QuantExporter,
|
||||
MXFP8QuantExporter,
|
||||
NVFP4QuantExporter,
|
||||
)
|
||||
from modelopt.onnx.export.nvfp4_exporter import _cast_fp4
|
||||
from modelopt.onnx.quantization.gs_patching import _export_value_info_proto
|
||||
from modelopt.onnx.quantization.qdq_utils import (
|
||||
_cast_fp8,
|
||||
apply_column_major_transformation,
|
||||
@@ -484,6 +492,57 @@ class TestCastFunctions:
|
||||
assert np.all(result == expected_array)
|
||||
|
||||
|
||||
class TestFP8QuantExporter:
|
||||
"""Test suite for FP8QuantExporter."""
|
||||
|
||||
def test_bf16_weights_and_scale_are_compressed(self):
|
||||
weight_data = np.array([0.001312255859375], dtype=ml_dtypes.bfloat16)
|
||||
scale_data = np.array(0.00099945068359375, dtype=ml_dtypes.bfloat16)
|
||||
weight = gs.Constant("weight", weight_data)
|
||||
scale = gs.Constant("linear/weight_quantizer/scale", scale_data)
|
||||
quantized = gs.Variable("quantized", dtype=np.uint8, shape=weight_data.shape)
|
||||
dequantized = gs.Variable(
|
||||
"dequantized", dtype=TensorProto.BFLOAT16, shape=weight_data.shape
|
||||
)
|
||||
value_info = _export_value_info_proto(dequantized, do_type_check=True)
|
||||
assert value_info.type.tensor_type.elem_type == TensorProto.BFLOAT16
|
||||
graph = gs.Graph(
|
||||
nodes=[
|
||||
gs.Node(
|
||||
op="TRT_FP8QuantizeLinear",
|
||||
inputs=[weight, scale],
|
||||
outputs=[quantized],
|
||||
),
|
||||
gs.Node(
|
||||
op="TRT_FP8DequantizeLinear",
|
||||
inputs=[quantized, scale],
|
||||
outputs=[dequantized],
|
||||
),
|
||||
],
|
||||
outputs=[dequantized],
|
||||
opset=23,
|
||||
)
|
||||
|
||||
converted_model = FP8QuantExporter.compress_weights(gs.export_onnx(graph))
|
||||
|
||||
onnx.checker.check_model(converted_model)
|
||||
assert [node.op_type for node in converted_model.graph.node] == ["DequantizeLinear"]
|
||||
assert converted_model.graph.output[0].type.tensor_type.elem_type == TensorProto.BFLOAT16
|
||||
fp8_weight = next(
|
||||
initializer
|
||||
for initializer in converted_model.graph.initializer
|
||||
if initializer.name == "linear/weight_quantizer/fp8_weights"
|
||||
)
|
||||
assert fp8_weight.data_type == TensorProto.FLOAT8E4M3FN
|
||||
assert fp8_weight.raw_data == b"\x3b"
|
||||
output_scale = next(
|
||||
initializer
|
||||
for initializer in converted_model.graph.initializer
|
||||
if initializer.name == scale.name
|
||||
)
|
||||
assert output_scale.data_type == TensorProto.BFLOAT16
|
||||
|
||||
|
||||
class TestMXFP8QuantExporter:
|
||||
"""Test suite for MXFP8QuantExporter."""
|
||||
|
||||
|
||||
@@ -13,7 +13,9 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import copy
|
||||
import json
|
||||
import tempfile
|
||||
from contextlib import nullcontext
|
||||
|
||||
import numpy as np
|
||||
@@ -24,6 +26,7 @@ import torch
|
||||
import torch.nn as nn
|
||||
from _test_utils.torch.deploy.lib_test_models import BaseDeployModel, get_deploy_models
|
||||
|
||||
import modelopt.torch.quantization as mtq
|
||||
from modelopt.onnx.utils import get_batch_size_from_bytes, validate_batch_size
|
||||
from modelopt.torch._deploy.utils import (
|
||||
OnnxBytes,
|
||||
@@ -57,6 +60,64 @@ deploy_benchmark_dynamo = {
|
||||
}
|
||||
|
||||
|
||||
class _FP8ModelWithBuffer(nn.Sequential):
|
||||
fp32_buffer: torch.Tensor
|
||||
|
||||
def forward(self, inputs):
|
||||
return super().forward(inputs) + self.fp32_buffer
|
||||
|
||||
|
||||
def _make_fp8_model(source_dtype, kind="fp8"):
|
||||
if kind == "format":
|
||||
model = nn.Sequential(*(nn.Linear(128, 128, bias=False) for _ in range(2)))
|
||||
sample_input = torch.ones(1, 128, dtype=source_dtype)
|
||||
config = copy.deepcopy(mtq.FP8_DEFAULT_CFG)
|
||||
config["quant_cfg"].extend(
|
||||
[
|
||||
{
|
||||
"quantizer_name": "1.weight_quantizer",
|
||||
"cfg": {"num_bits": 4, "block_sizes": {-1: 128, "type": "static"}},
|
||||
},
|
||||
{"quantizer_name": "1.input_quantizer", "enable": False},
|
||||
]
|
||||
)
|
||||
else:
|
||||
model = _FP8ModelWithBuffer(
|
||||
nn.Conv2d(1, 1, 1, bias=False),
|
||||
nn.Flatten(),
|
||||
nn.Linear(4, 4, bias=False),
|
||||
)
|
||||
sample_input = torch.ones(1, 1, 2, 2, dtype=source_dtype)
|
||||
config = mtq.FP8_DEFAULT_CFG
|
||||
model = model.eval().to(source_dtype)
|
||||
if kind != "format":
|
||||
model.register_buffer("fp32_buffer", torch.ones(4))
|
||||
if kind == "parameters":
|
||||
model.register_parameter("unused_fp32_parameter", nn.Parameter(torch.ones(1)))
|
||||
quantized_model = mtq.quantize(
|
||||
model,
|
||||
config,
|
||||
forward_loop=lambda quantized_model: quantized_model(sample_input),
|
||||
)
|
||||
if kind != "format":
|
||||
# Keep calibration nonzero while exercising Conv scale underflow during export.
|
||||
with torch.no_grad():
|
||||
quantized_model[0].weight.fill_(1e-38 if source_dtype == torch.bfloat16 else 1.0)
|
||||
return quantized_model, sample_input
|
||||
|
||||
|
||||
def _export_fp8_model(source_dtype, weights_dtype):
|
||||
model, sample_input = _make_fp8_model(source_dtype)
|
||||
onnx_bytes, _ = get_onnx_bytes_and_metadata(
|
||||
model,
|
||||
(sample_input,),
|
||||
weights_dtype=weights_dtype,
|
||||
onnx_opset=23,
|
||||
)
|
||||
onnx_bytes_obj = OnnxBytes.from_bytes(onnx_bytes)
|
||||
return onnx.load_model_from_string(onnx_bytes_obj.get_onnx_model_file_bytes())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model", deploy_benchmark_dynamo.values(), ids=deploy_benchmark_dynamo.keys()
|
||||
)
|
||||
@@ -157,6 +218,87 @@ def test_onnx_export_and_inputs(model: BaseDeployModel):
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("source_dtype", "weights_dtype", "expected_onnx_dtype"),
|
||||
[
|
||||
(torch.bfloat16, "bf16", onnx.TensorProto.BFLOAT16),
|
||||
(torch.float32, "fp16", onnx.TensorProto.FLOAT16),
|
||||
],
|
||||
ids=["bf16-weight-focused-buffer", "fp32-to-fp16"],
|
||||
)
|
||||
def test_fp8_export_with_supported_weights_dtype(source_dtype, weights_dtype, expected_onnx_dtype):
|
||||
exported_model = _export_fp8_model(source_dtype, weights_dtype)
|
||||
|
||||
onnx.checker.check_model(exported_model, full_check=True)
|
||||
assert {"TRT_FP8QuantizeLinear", "TRT_FP8DequantizeLinear"}.isdisjoint(
|
||||
node.op_type for node in exported_model.graph.node
|
||||
)
|
||||
initializer_by_name = {
|
||||
initializer.name: initializer for initializer in exported_model.graph.initializer
|
||||
}
|
||||
fp8_weight_dq_nodes = [
|
||||
node
|
||||
for node in exported_model.graph.node
|
||||
if node.op_type == "DequantizeLinear"
|
||||
and node.input[0] in initializer_by_name
|
||||
and initializer_by_name[node.input[0]].data_type == onnx.TensorProto.FLOAT8E4M3FN
|
||||
]
|
||||
assert len(fp8_weight_dq_nodes) == 2
|
||||
for node in fp8_weight_dq_nodes:
|
||||
assert initializer_by_name[node.input[1]].data_type == expected_onnx_dtype
|
||||
assert {0x7F, 0xFF}.isdisjoint(initializer_by_name[node.input[0]].raw_data)
|
||||
assert all(
|
||||
value.type.tensor_type.elem_type == expected_onnx_dtype
|
||||
for value in exported_model.graph.input
|
||||
)
|
||||
expected_output_dtype = (
|
||||
onnx.TensorProto.FLOAT if weights_dtype == "bf16" else expected_onnx_dtype
|
||||
)
|
||||
assert all(
|
||||
value.type.tensor_type.elem_type == expected_output_dtype
|
||||
for value in exported_model.graph.output
|
||||
)
|
||||
if weights_dtype == "bf16":
|
||||
fp32_buffer = next(
|
||||
initializer
|
||||
for initializer in exported_model.graph.initializer
|
||||
if initializer.name.endswith("fp32_buffer")
|
||||
)
|
||||
assert fp32_buffer.data_type == onnx.TensorProto.FLOAT
|
||||
assert any(
|
||||
node.op_type == "Add" and fp32_buffer.name in node.input
|
||||
for node in exported_model.graph.node
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("kind", "source_dtype", "weights_dtype", "error"),
|
||||
[
|
||||
("fp8", torch.float32, "bf16", "torch.float32"),
|
||||
("fp8", torch.bfloat16, "fp16", "torch.bfloat16"),
|
||||
("parameters", torch.bfloat16, "bf16", "torch.bfloat16, torch.float32"),
|
||||
("format", torch.bfloat16, "bf16", "torch.bfloat16"),
|
||||
],
|
||||
ids=["fp32-to-bf16", "bf16-to-fp16", "mixed-parameters", "mixed-format"],
|
||||
)
|
||||
def test_fp8_export_rejects_unsupported_dtype_conversion(
|
||||
kind, source_dtype, weights_dtype, error, monkeypatch, tmp_path
|
||||
):
|
||||
monkeypatch.setattr(tempfile, "tempdir", str(tmp_path))
|
||||
model, sample_input = _make_fp8_model(source_dtype, kind)
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match=rf"Converting .* to {weights_dtype.upper()}.*source parameter dtypes: {error}",
|
||||
):
|
||||
get_onnx_bytes_and_metadata(
|
||||
model,
|
||||
(sample_input,),
|
||||
weights_dtype=weights_dtype,
|
||||
onnx_opset=23,
|
||||
)
|
||||
assert not any(tmp_path.iterdir())
|
||||
|
||||
|
||||
class SingleArgModel(nn.Module):
|
||||
def forward(self, x: torch.Tensor):
|
||||
return torch.add(x, x) - x
|
||||
|
||||
Reference in New Issue
Block a user