[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:
Ajinkya Rasane
2026-09-10 18:30:26 +00:00
committed by GitHub
co-authored by Codex Codex
parent d19925e446
commit 7f7c46d820
6 changed files with 286 additions and 35 deletions
+1
View File
@@ -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).
+31 -9
View File
@@ -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",
+9 -3
View File
@@ -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)
+43 -22
View File
@@ -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)
+60 -1
View File
@@ -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