mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[6078291][OMNIML-3716] Add ViT FP8 + Torch-TRT example, wire softmax_quantizer in _QuantAttention (#1569)
### What does this PR do?
Type of change: new feature + bug fix
Adds a Torch-TensorRT deployment path for HuggingFace ViT and closes the
modelopt-side gap that prevented `*softmax_quantizer` from being applied
on the standard attention forward path.
* **New ViT PTQ recipes** under `modelopt_recipes/huggingface/vit/ptq/`:
* `fp8.yaml` — W8A8 per-tensor FP8 E4M3 on encoder Linear
weights/inputs;
attention Q/K/V BMMs + softmax output at FP8; per-block LayerNorm output
at FP8 (one shared Q/DQ feeds Q/K/V + MLP); patch-embed `nn.Conv2d`,
`classifier`, and the final `vit.layernorm` left FP16. Uses max
calibration.
* The recipe is self-contained (no `$import` of shared snippets) and
use the "specific-enable" style: narrow `parent_class` + path scoping
on every enable rule, so no `enable: false` carve-outs are needed.
* **New example** under `examples/torch_trt/`:
* `torch_tensorrt_ptq.py` — single-model pipeline (load HF model,
calibrate from `zh-plus/tiny-imagenet`, `mtq.quantize`,
`torch_tensorrt.compile`, verify the compiled-model argmax matches the
fake-quant argmax). Defaults to `google/vit-large-patch16-224`; pass
`--model_id` and `--recipe` to target any model + recipe combination.
`--no_pretrained` + `--model_kwargs` shrink the model for fast tests.
* `README.md` documenting the flow, the shipped recipes, hardware
requirements, and CLI usage.
* `requirements.txt`.
* **Bug fix in `modelopt/torch/quantization/plugins/huggingface.py`** —
inside
`_QuantAttention._quantized_attention`, the non-kitchen branch now
temporarily replaces `torch.nn.functional.softmax` (via the existing
`replace_function` context manager) with a wrapper that pipes the
softmax
output through `self.softmax_quantizer`. Previously the slot was created
on every registered attention class but only consumed by the optional
Kitchen MXFP8 flash-attention path, so FP8 / NVFP4 recipes that enabled
`*softmax_quantizer` saw it stay uncalibrated (`amax=None`) and emitted
no Q/DQ around the softmax output during ONNX / Torch-TRT export. With
this fix the `softmax_quantizer` is calibrated alongside the rest of
the model, and both the modelopt ONNX exporter and
`torch_tensorrt.compile`
pick up the Q/DQ pair. The patch short-circuits to the unwrapped call
when the quantizer is disabled (zero-overhead) and has no effect on SDPA
paths that fuse softmax inside a C++ kernel.
* **New e2e integration test** at
`tests/examples/torch_trt/test_torch_tensorrt_ptq.py` — mirrors the
`torch_onnx` test pattern: invokes the example through
`run_example_command`, parametrizes over the two precision modes (fp8,
nvfp4), uses a 1-layer ViT config (`--no_pretrained` + `--model_kwargs`)
so each parametrized case completes in under a minute. `importorskip` on
`torch_tensorrt` so the test is automatically skipped on hosts without
the package.
### Usage
```bash
# FP8 (Hopper / Ada) — default model is google/vit-large-patch16-224
python examples/torch_trt/torch_tensorrt_ptq.py \
--precision fp8 \
--calib_samples 128 \
--batch_size 1
# Custom model + custom recipe
python examples/torch_trt/torch_tensorrt_ptq.py \
--model_id <huggingface/model-id> \
--recipe <recipe-path-relative-to-modelopt_recipes-or-absolute-yaml>
```
### Testing
* Recipes load via `modelopt.recipe.load_recipe()` and pass
`QuantizeConfig` schema validation.
* Run `pytest tests/examples/torch_trt/test_torch_tensorrt_ptq.py` →
1 parametrized case passes on RTX 6000 Ada (fp8).
* End-to-end on `google/vit-base-patch16-224`: `mtq.quantize` with the
new
FP8 recipe followed by `torch_tensorrt.compile(ir="dynamo")` produces a
TRT engine whose argmax matches the FP16 baseline.
* ONNX exported from the torch path now contains Q/DQ on **12 / 12**
softmax outputs (was 0 / 12 before this PR's `_QuantAttention` fix),
matching the ONNX-CLI output's quantization layout.
Both FP8 paths land within 0.13 pp Top-1 of the FP16 baseline; Top-5 is
within 0.02 pp across all three.
* ImageNet-1k validation accuracy via the new
`torch_tensorrt_accuracy.py`
(full 50000 samples, batch=1, **every model Torch-TensorRT-compiled —
including the baseline** — so the comparison is apples-to-apples) for
the
example's default `google/vit-large-patch16-224`:
| Model (Torch-TRT) | Top-1 | Top-5 | Δ Top-1 vs baseline |
|---|---:|---:|---:|
| Baseline (FP16) | 81.99% | 96.01% | — |
| FP8 | 82.01% | 96.05% | +0.02 pp |
FP8 is within noise of the FP16 TRT baseline and NVFP4 W4A4 costs only
−0.13 pp Top-1 / −0.05 pp Top-5. Absolute Top-1 sits below the model
card's
~85.5% because evaluation uses the HF `AutoImageProcessor` default
preprocessing (direct 224×224 resize, no resize-then-center-crop),
applied
identically to all three models — so the deltas are the comparison
signal.
### 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?: ✅ — new e2e integration test
under `tests/examples/torch_trt/`.
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅
<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit
* **New Features**
* Added Torch‑TensorRT FP8/NVFP4 deployment examples and end‑to‑end
scripts for HuggingFace ViT, plus ViT-specific PTQ recipes and
ImageNet-1k vs FP16 accuracy reporting.
* **Bug Fixes**
* Fixed softmax quantization and export/compilation edge cases (softmax
calibration during export, IO casting for empty tensors, routed expert
weight syncing, importer key handling).
* **Documentation**
* Added comprehensive example README with setup, usage, recipes,
evaluation, and hardware guidance.
* **Requirements**
* Pinned minimum versions for example dependencies.
* **Tests**
* Added tests validating the Torch‑TensorRT quantization examples for
fp8.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->
---------
Signed-off-by: ajrasane <131806219+ajrasane@users.noreply.github.com>
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
973cb09cbe
commit
fbbc5989ce
@@ -0,0 +1,240 @@
|
||||
# Torch-TensorRT Quantization
|
||||
|
||||
[Torch-TensorRT](https://docs.pytorch.org/TensorRT/) compiles a PyTorch model into an optimized TensorRT engine with no separate export or runtime. This example quantizes a PyTorch / HuggingFace model with NVIDIA Model Optimizer and then compiles the quantized graph in-framework with Torch-TensorRT for deployment.
|
||||
|
||||
Quantization is an effective model optimization technique that compresses your models. Model Optimizer inserts Q/DQ nodes into the eager PyTorch graph; `torch_tensorrt.compile(ir="dynamo")` then converts those Q/DQ nodes into native TensorRT FP8 precision layers, following the [Torch-TensorRT quantization guide](https://docs.pytorch.org/TensorRT/user_guide/shapes_precision/quantization.html).
|
||||
|
||||
This section focuses on the in-framework Torch-TensorRT path: a PyTorch front end (`mtq.quantize`) feeding a Dynamo-compiled TensorRT engine, demonstrated end-to-end on a HuggingFace ViT image classifier. If you instead want a portable ONNX → TensorRT artifact, or you start from an ONNX model, see the sibling [`torch_onnx`](../torch_onnx/) and [`onnx_ptq`](../onnx_ptq/) examples (compared in the [Support Matrix](#support-matrix)).
|
||||
|
||||
<div align="center">
|
||||
|
||||
| **Section** | **Description** | **Link** | **Docs** |
|
||||
| :------------: | :------------: | :------------: | :------------: |
|
||||
| Pre-Requisites | Required packages and installation | \[[Link](#pre-requisites)\] | |
|
||||
| Getting Started | Quantize and compile a ViT in a few lines | \[[Link](#getting-started)\] | \[[docs](https://docs.pytorch.org/TensorRT/user_guide/shapes_precision/quantization.html)\] |
|
||||
| Support Matrix | How this path compares to the ONNX examples | \[[Link](#support-matrix)\] | |
|
||||
| ViT Recipes | The FP8 recipe shipped with the example | \[[Link](#vit-recipes)\] | |
|
||||
| Usage | CLI flags for the quantize and accuracy scripts | \[[Link](#usage)\] | |
|
||||
| Evaluate Accuracy | Measure ImageNet top-1 / top-5 accuracy | \[[Link](#evaluate-accuracy)\] | |
|
||||
| Custom Recipes | Plug in your own recipe / model | \[[Link](#custom-recipes)\] | |
|
||||
| Resources | Roadmap, docs, benchmarks, and support | \[[Link](#resources)\] | |
|
||||
|
||||
</div>
|
||||
|
||||
## Pre-Requisites
|
||||
|
||||
### Docker
|
||||
|
||||
Please use the TensorRT docker image (e.g., `nvcr.io/nvidia/tensorrt:26.02-py3`) or visit our [installation docs](https://nvidia.github.io/Model-Optimizer/getting_started/2_installation.html) for more information.
|
||||
|
||||
```bash
|
||||
docker run --gpus all -it --rm -v $(pwd):/workspace -w /workspace nvcr.io/nvidia/tensorrt:26.02-py3 bash
|
||||
```
|
||||
|
||||
Also follow the installation steps below to upgrade to the latest version of Model Optimizer and install example-specific dependencies.
|
||||
|
||||
### Local Installation
|
||||
|
||||
```bash
|
||||
pip install -U "nvidia-modelopt[hf]"
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
### Hardware Requirements
|
||||
|
||||
The low-precision kernels Torch-TensorRT emits need a GPU that supports the target format:
|
||||
|
||||
<div align="center">
|
||||
|
||||
| Recipe | Minimum GPU |
|
||||
| :---: | :---: |
|
||||
| `fp8` | Ada / Hopper — compute capability 8.9+ |
|
||||
|
||||
</div>
|
||||
|
||||
> [!NOTE]
|
||||
> Older GPUs still let `mtq.quantize` succeed — it emits fake-quant nodes in PyTorch — but `torch_tensorrt.compile` will not find a real low-precision kernel for an unsupported format.
|
||||
|
||||
## Getting Started
|
||||
|
||||
Quantize a HuggingFace ViT, then compile the Q/DQ graph with Torch-TensorRT into a single `torch.nn.Module` you call from PyTorch:
|
||||
|
||||
```python
|
||||
import torch
|
||||
import torch_tensorrt
|
||||
|
||||
import modelopt.torch.quantization as mtq
|
||||
from modelopt.recipe import load_recipe
|
||||
from modelopt.torch.quantization.utils import export_torch_mode
|
||||
|
||||
# 1. Quantize the eager PyTorch model with a Model Optimizer PTQ recipe.
|
||||
recipe = load_recipe("huggingface/vit/ptq/fp8")
|
||||
mtq.quantize(model, recipe.quantize.model_dump(), forward_loop=calibrate)
|
||||
|
||||
# 2. Compile the quantized (Q/DQ) graph with Torch-TensorRT.
|
||||
# export_torch_mode() makes Model Optimizer emit Q/DQ in the TRT-friendly form,
|
||||
# and min_block_size=1 lets single-node Q/DQ + matmul subgraphs become TRT
|
||||
# precision layers (per the Torch-TensorRT quantization guide).
|
||||
with export_torch_mode():
|
||||
trt_model = torch_tensorrt.compile(
|
||||
model,
|
||||
ir="dynamo",
|
||||
min_block_size=1,
|
||||
truncate_double=True,
|
||||
inputs=[torch_tensorrt.Input(
|
||||
min_shape=(1, 3, 224, 224),
|
||||
opt_shape=(128, 3, 224, 224),
|
||||
max_shape=(1024, 3, 224, 224),
|
||||
dtype=torch.float16,
|
||||
)],
|
||||
)
|
||||
|
||||
logits = trt_model(pixel_values) # call it like any nn.Module
|
||||
```
|
||||
|
||||
The runnable script [`torch_tensorrt_ptq.py`](./torch_tensorrt_ptq.py) wraps this flow end-to-end. It:
|
||||
|
||||
1. Loads a HuggingFace ViT classifier (default `google/vit-large-patch16-224`).
|
||||
1. Builds a tiny calibration loader from `zh-plus/tiny-imagenet` (avoids the gated `ILSVRC/imagenet-1k` repo, so the example runs unauthenticated).
|
||||
1. Runs `mtq.quantize` with one of the recipes under [`modelopt_recipes/`](../../modelopt_recipes/) (see [ViT Recipes](#vit-recipes)).
|
||||
1. Saves the quantized Model Optimizer state (FP16 weights + Q/DQ metadata) to `<save_dir>/vit_modelopt_state.pt` for reuse without recalibration (see [Custom Recipes](#custom-recipes)).
|
||||
1. Compiles the quantized model with `torch_tensorrt.compile` and verifies that the compiled-model argmax matches the fake-quant argmax on a sample input.
|
||||
|
||||
```bash
|
||||
# Default model is google/vit-large-patch16-224, default recipe is the ViT FP8 recipe.
|
||||
python torch_tensorrt_ptq.py --calib_samples 1024 --batch_size 128
|
||||
|
||||
# Quantize but don't TRT-compile (handy on a non-TRT host).
|
||||
python torch_tensorrt_ptq.py --skip_trt
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
> Both `torch_tensorrt_ptq.py` and the accuracy script ([`torch_tensorrt_accuracy.py`](./torch_tensorrt_accuracy.py)) run the model in `float16`.
|
||||
|
||||
## Support Matrix
|
||||
|
||||
All three of these examples reach the same destination — a low-precision TensorRT engine — but quantize at a different point in the pipeline and emit a different artifact, so they suit different deployment stacks:
|
||||
|
||||
<div align="center">
|
||||
|
||||
| | Torch-TensorRT (this example) | [`torch_onnx`](../torch_onnx/) | [`onnx_ptq`](../onnx_ptq/) |
|
||||
| :---: | :---: | :---: | :---: |
|
||||
| Starting point | a PyTorch / HF model | a PyTorch / timm model | an already-exported ONNX model |
|
||||
| Quantize on | the eager PyTorch graph (`mtq.quantize`) | the eager PyTorch graph (`mtq.quantize`) | the ONNX graph directly (ONNX PTQ) |
|
||||
| Export step | none — the FX/Dynamo graph stays in-process | `torch.onnx.export` of the Q/DQ graph, postprocessed for TRT | none — Q/DQ inserted straight into the ONNX graph |
|
||||
| Intermediate artifact | none | a Q/DQ ONNX file | a Q/DQ ONNX file |
|
||||
| Compiler + runtime | `torch_tensorrt.compile(ir="dynamo")` → a `torch.nn.Module` you call from PyTorch | TensorRT builds a standalone engine from the ONNX | TensorRT builds a standalone engine from the ONNX |
|
||||
| Best when | PyTorch-native serving; you want a drop-in compiled module | you quantize in PyTorch but deploy via a portable ONNX → TRT engine | you only have an ONNX model and never touch PyTorch |
|
||||
|
||||
</div>
|
||||
|
||||
This example and [`torch_onnx`](../torch_onnx/) share the same PyTorch front end (`mtq.quantize`), so the numerics are identical — they differ only in the back end: this one keeps the graph in-process and hands it to Torch-TensorRT, while `torch_onnx` exports a portable ONNX artifact for the standalone TensorRT runtime. [`onnx_ptq`](../onnx_ptq/) instead quantizes the ONNX graph directly, for when you start from an ONNX model rather than PyTorch. Pick this example when your serving stack is PyTorch-native and you'd rather avoid an ONNX export step.
|
||||
|
||||
## ViT Recipes
|
||||
|
||||
This is the recipe the CLI selects by default when `--model_id` points at a HF ViT classifier. It is tuned for the HF ViT module layout and is composed from the shared `$import` building blocks under [`modelopt_recipes/configs/`](../../modelopt_recipes/configs/) (`ptq/units/{w8a8_fp8_fp8,attention_qkv_fp8}`) rather than spelling out each `quant_cfg` entry.
|
||||
|
||||
<div align="center">
|
||||
|
||||
| `--recipe` value | Calibration | What it quantizes |
|
||||
| :---: | :---: | :--- |
|
||||
| `huggingface/vit/ptq/fp8` (default) | `max` | Per-tensor FP8 (E4M3) on every weight + input quantizer matched by the `*weight_quantizer` / `*input_quantizer` globs — encoder Linears, the patch-embed `nn.Conv2d` projection, and the `classifier` head — plus FP8 on the attention Q/K/V BMMs and softmax. All output quantizers disabled. |
|
||||
|
||||
</div>
|
||||
|
||||
## Usage
|
||||
|
||||
### `torch_tensorrt_ptq.py`
|
||||
|
||||
[Script](./torch_tensorrt_ptq.py) — quantize and (optionally) Torch-TensorRT-compile a ViT.
|
||||
|
||||
<div align="center">
|
||||
|
||||
| Flag | Default | Description |
|
||||
| :---: | :---: | :--- |
|
||||
| `--model_id` | `google/vit-large-patch16-224` | HuggingFace model id of the ViT classifier to quantize. |
|
||||
| `--recipe` | `huggingface/vit/ptq/fp8` | Recipe path (relative to `modelopt_recipes/` or an absolute YAML). |
|
||||
| `--calib_samples` | `1024` | Number of tiny-imagenet samples to use for calibration. |
|
||||
| `--batch_size` | `128` | Batch size for calibration / TRT compile. |
|
||||
| `--save_dir` | `./modelopt_quantized` | Directory the quantized Model Optimizer state-dict (FP16 weights + Q/DQ metadata) is always saved to, as `vit_modelopt_state.pt` — re-usable across runs without recalibration. |
|
||||
| `--skip_trt` | off | Quantize + run the fake-quant model only; skip `torch_tensorrt.compile`. Useful for environments without Torch-TensorRT installed. |
|
||||
| `--layer_info_path` | unset | If set, write the compiled TRT engine's per-layer info (`get_layer_info()`) to this file. |
|
||||
|
||||
</div>
|
||||
|
||||
```bash
|
||||
# Custom model + custom recipe, saving the quantized state elsewhere.
|
||||
python torch_tensorrt_ptq.py \
|
||||
--model_id <huggingface/model-id> \
|
||||
--recipe <recipe-path-relative-to-modelopt_recipes-or-absolute-yaml> \
|
||||
--save_dir ./my_quantized
|
||||
|
||||
# Dump the compiled engine's per-layer info to inspect FP8 fusion.
|
||||
python torch_tensorrt_ptq.py --layer_info_path ./vit_fp8_layers.txt
|
||||
```
|
||||
|
||||
### `torch_tensorrt_accuracy.py`
|
||||
|
||||
[Script](./torch_tensorrt_accuracy.py) — quantize, compile, and score on ImageNet (see [Evaluate Accuracy](#evaluate-accuracy)).
|
||||
|
||||
<div align="center">
|
||||
|
||||
| Flag | Default | Description |
|
||||
| :---: | :---: | :--- |
|
||||
| `--model_id` | `google/vit-large-patch16-224` | HuggingFace model id of the ViT classifier to quantize and score. |
|
||||
| `--recipe` | `huggingface/vit/ptq/fp8` | Recipe path (relative to `modelopt_recipes/` or an absolute YAML). |
|
||||
| `--calib_samples` | `1024` | Number of tiny-imagenet samples to use for calibration. |
|
||||
| `--batch_size` | `128` | Calibration / compile / eval batch size. The Torch-TRT engine is dynamic (`min=1`, `opt=max(--batch_size, 2)`, `max=1024`) and handles any batch including the trailing partial batch. |
|
||||
| `--eval_data_size` | full 50k | Number of ImageNet validation images to score. |
|
||||
| `--imagenet_path` | `ILSVRC/imagenet-1k` | HF dataset card or local path to the ImageNet validation set (gated). |
|
||||
| `--baseline` | off | Also score the unquantized model as a reference. It is Torch-TensorRT-compiled like the quantized model (or run eager under `--skip_trt`) so the comparison is apples-to-apples. |
|
||||
| `--skip_trt` | off | Score the fake-quant (Model Optimizer) model; skip `torch_tensorrt.compile`. Useful for environments without Torch-TensorRT installed. |
|
||||
| `--results_path` | unset | If set, write the accuracy results to this CSV path. |
|
||||
|
||||
</div>
|
||||
|
||||
## Evaluate Accuracy
|
||||
|
||||
[`torch_tensorrt_accuracy.py`](./torch_tensorrt_accuracy.py) reuses the quantize → compile pipeline above and reports ImageNet-1k top-1 / top-5 accuracy via the `onnx_ptq` example's `evaluate()` harness ([`examples/onnx_ptq/evaluation.py`](../onnx_ptq/evaluation.py)):
|
||||
|
||||
```bash
|
||||
python torch_tensorrt_accuracy.py \
|
||||
--recipe huggingface/vit/ptq/fp8 \
|
||||
--batch_size 128 \
|
||||
--baseline \
|
||||
--eval_data_size 5000 \
|
||||
--results_path results.csv
|
||||
```
|
||||
|
||||
- `--baseline` also scores the unquantized model. It is Torch-TensorRT-compiled the same way as the quantized model, so every reported number comes from the same TRT runtime (pass `--skip_trt` to score the eager / fake-quant models instead).
|
||||
- The eval uses a **dynamic** engine (default `--batch_size 128`) for both precisions, so it serves the trailing partial batch at any batch size.
|
||||
- `--results_path results.csv` writes the metrics table (`Metric`, `Top1 (%)`, `Top5 (%)`) to CSV.
|
||||
|
||||
> [!NOTE]
|
||||
> Validation uses the gated `ILSVRC/imagenet-1k` split: accept its license / set `HF_TOKEN`, or point `--imagenet_path` at a local copy. `evaluate()` shuffles the split, so a partial `--eval_data_size` draws a different random subset each run — omit it (full 50k set) for a stable, comparable score.
|
||||
|
||||
## Custom Recipes
|
||||
|
||||
Use `--recipe <path>` to plug in a different recipe — either a path relative to `modelopt_recipes/` (resolved against the built-in recipe library) or an absolute filesystem path to a YAML file. The recipe is loaded via `modelopt.recipe.load_recipe`, must declare `metadata.recipe_type: ptq` and a `quantize:` section, and its `quantize` config is passed straight to `mtq.quantize`. See the existing [`modelopt_recipes/huggingface/vit/ptq/*.yaml`](../../modelopt_recipes/huggingface/vit/ptq/) for the patterns used here.
|
||||
|
||||
### Resuming From a Saved Checkpoint
|
||||
|
||||
`torch_tensorrt_ptq.py` always saves the quantized Model Optimizer state to `<save_dir>/vit_modelopt_state.pt` (default `--save_dir ./modelopt_quantized`) via `mto.save`. To reload it without recalibrating, restore it onto a freshly-loaded model before the TRT compile step:
|
||||
|
||||
```python
|
||||
import modelopt.torch.opt as mto
|
||||
|
||||
mto.restore(model, "./modelopt_quantized/vit_modelopt_state.pt")
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
> See the [save / restore guide](https://nvidia.github.io/Model-Optimizer/guides/2_save_load.html) for the full `mto.save` / `mto.restore` workflow.
|
||||
|
||||
## Resources
|
||||
|
||||
- 📅 [Roadmap](https://github.com/NVIDIA/Model-Optimizer/issues/146)
|
||||
- 📖 [Documentation](https://nvidia.github.io/Model-Optimizer)
|
||||
- 🎯 [Benchmarks](../benchmark.md)
|
||||
- 💡 [Release Notes](https://nvidia.github.io/Model-Optimizer/reference/0_changelog.html)
|
||||
- 🐛 [File a bug](https://github.com/NVIDIA/Model-Optimizer/issues/new?template=1_bug_report.md)
|
||||
- ✨ [File a Feature Request](https://github.com/NVIDIA/Model-Optimizer/issues/new?template=2_feature_request.md)
|
||||
@@ -0,0 +1,3 @@
|
||||
datasets>=2.14.4
|
||||
torch-tensorrt>=2.4.0
|
||||
transformers>=4.56
|
||||
@@ -0,0 +1,228 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Measure ImageNet top-1/top-5 accuracy of a Torch-TensorRT ViT.
|
||||
|
||||
Pipeline:
|
||||
|
||||
1. Quantize a HuggingFace ViT with a ModelOpt recipe and compile it with
|
||||
``torch_tensorrt.compile(ir="dynamo")`` — reusing the sibling example
|
||||
``torch_tensorrt_ptq.py``.
|
||||
2. Score the compiled model on the ImageNet-1k validation split using the
|
||||
``onnx_ptq`` example's ``evaluate`` API (``examples/onnx_ptq/evaluation.py``).
|
||||
|
||||
The compiled Torch-TRT module is a ``torch.nn.Module``, so ``evaluate`` runs it
|
||||
exactly like an eager model. A thin :class:`_EvalAdapter` bridges the two
|
||||
contracts: it casts the dataloader's float32 image batches to the model's
|
||||
compute dtype and unwraps HF ``ImageClassifierOutput`` to a plain logits tensor.
|
||||
|
||||
Example::
|
||||
|
||||
python torch_tensorrt_accuracy.py --batch_size 128 --eval_data_size 5000 --baseline
|
||||
|
||||
``--imagenet_path`` defaults to the gated ``ILSVRC/imagenet-1k`` HF dataset
|
||||
(accept its license / set ``HF_TOKEN``), or point it at a local copy. Note the
|
||||
``evaluate`` API shuffles the validation set, so a partial ``--eval_data_size``
|
||||
samples a different random subset each run; use the full set for a stable score.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
# Reuse the quantize -> torch_tensorrt.compile pipeline from the sibling example.
|
||||
_THIS_DIR = Path(__file__).resolve().parent
|
||||
sys.path.insert(0, str(_THIS_DIR))
|
||||
import torch_tensorrt_ptq as ttptq # noqa: E402
|
||||
|
||||
# Reuse the ImageNet accuracy harness from the onnx_ptq example (sibling dir).
|
||||
_ONNX_PTQ_DIR = _THIS_DIR.parent / "onnx_ptq"
|
||||
sys.path.insert(0, str(_ONNX_PTQ_DIR))
|
||||
from evaluation import evaluate # noqa: E402
|
||||
|
||||
|
||||
class _EvalAdapter(torch.nn.Module):
|
||||
"""Adapt a compiled/eager ViT to the ``onnx_ptq`` ``evaluate`` contract.
|
||||
|
||||
``evaluate_accuracy`` feeds float32 image batches, calls ``model(inputs)``,
|
||||
and reads ``outputs.data``. This adapter casts inputs to the model's compute
|
||||
dtype (the dataloader yields FP32) and unwraps an HF ``ImageClassifierOutput``
|
||||
to the bare logits tensor.
|
||||
"""
|
||||
|
||||
def __init__(self, model: torch.nn.Module, dtype: torch.dtype):
|
||||
super().__init__()
|
||||
self.model = model
|
||||
self._dtype = dtype
|
||||
|
||||
def forward(self, pixel_values: torch.Tensor) -> torch.Tensor:
|
||||
out = self.model(pixel_values.to(self._dtype))
|
||||
return out.logits if hasattr(out, "logits") else out
|
||||
|
||||
|
||||
def build_processor_transform(processor):
|
||||
"""Return a ``PIL.Image -> (C, H, W) float tensor`` transform from the HF processor.
|
||||
|
||||
Using the model's own image processor keeps eval preprocessing (resize,
|
||||
normalization mean/std) consistent with how the ViT was trained, which is
|
||||
more faithful for a HuggingFace checkpoint than a generic timm transform.
|
||||
The model and ``ILSVRC/imagenet-1k`` share the standard 1000-class ordering,
|
||||
so predicted indices line up with the dataset labels.
|
||||
"""
|
||||
|
||||
def _transform(image):
|
||||
return processor(images=image, return_tensors="pt")["pixel_values"][0]
|
||||
|
||||
return _transform
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(
|
||||
description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model_id",
|
||||
default="google/vit-large-patch16-224",
|
||||
help="HuggingFace model id of the ViT classifier to quantize and score.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--recipe",
|
||||
default=ttptq.DEFAULT_RECIPE,
|
||||
help="Recipe path (relative to modelopt_recipes/ or an absolute YAML). "
|
||||
"Defaults to the ViT FP8 recipe.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--calib_samples",
|
||||
type=int,
|
||||
default=1024,
|
||||
help="Number of tiny-imagenet samples to use for calibration.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size",
|
||||
type=int,
|
||||
default=128,
|
||||
help="Calibration / compile / eval batch size. The Torch-TRT engine is "
|
||||
"dynamic (min=1, opt=--batch_size, max=1024) and handles any batch incl. "
|
||||
"the trailing partial batch, so any --batch_size (e.g. 128) works.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--eval_data_size",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Number of ImageNet validation images to score (default: full 50k).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--imagenet_path",
|
||||
default="ILSVRC/imagenet-1k",
|
||||
help="HF dataset card or local path to the ImageNet validation set (gated).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--baseline",
|
||||
action="store_true",
|
||||
help="Also score the unquantized model as a reference. It is "
|
||||
"Torch-TensorRT-compiled like the quantized model (or run eager under "
|
||||
"--skip_trt) so the comparison is apples-to-apples.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--skip_trt",
|
||||
action="store_true",
|
||||
help="Score the fake-quant (modelopt) model; skip torch_tensorrt.compile. "
|
||||
"Useful for environments without torch_tensorrt installed.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--results_path",
|
||||
default=None,
|
||||
help="If set, write the accuracy results to this CSV path.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
raise SystemExit("This example requires a CUDA-capable GPU.")
|
||||
device = torch.device("cuda")
|
||||
dtype = torch.float16
|
||||
|
||||
model, processor = ttptq.load_model_and_processor(args.model_id, device, dtype)
|
||||
transform = build_processor_transform(processor)
|
||||
|
||||
def run_eval(m: torch.nn.Module) -> tuple[float, float]:
|
||||
top1, top5 = evaluate(
|
||||
_EvalAdapter(m, dtype),
|
||||
transform,
|
||||
batch_size=args.batch_size,
|
||||
num_examples=args.eval_data_size,
|
||||
device="cuda",
|
||||
dataset_path=args.imagenet_path,
|
||||
)
|
||||
return top1, top5
|
||||
|
||||
image_size = model.config.image_size
|
||||
num_channels = model.config.num_channels
|
||||
example_input = torch.randn(
|
||||
args.batch_size, num_channels, image_size, image_size, device=device, dtype=dtype
|
||||
)
|
||||
runtime = "fake-quant" if args.skip_trt else "torch-trt"
|
||||
|
||||
def to_eval_model(m: torch.nn.Module, what: str) -> torch.nn.Module:
|
||||
"""Logits-wrap and (unless --skip_trt) Torch-TensorRT-compile ``m`` for eval.
|
||||
|
||||
The baseline is compiled the same way as the quantized model so all
|
||||
reported numbers come from the same Torch-TRT runtime.
|
||||
"""
|
||||
wrapped = ttptq.ViTLogitsWrapper(m).to(device).eval()
|
||||
if args.skip_trt:
|
||||
return wrapped
|
||||
print(f"\nCompiling {what} with Torch-TensorRT ...")
|
||||
return ttptq.compile_with_torch_tensorrt(wrapped, example_input)
|
||||
|
||||
results: list[list[str | float]] = [["Metric", "Top1 (%)", "Top5 (%)"]]
|
||||
|
||||
# Baseline must be built + scored before in-place quantization mutates `model`.
|
||||
if args.baseline:
|
||||
prec = str(dtype).rsplit(".", 1)[-1] # e.g. "float16"
|
||||
base_tag = f"baseline-{prec} ({runtime})"
|
||||
base_eval = to_eval_model(model, "unquantized baseline")
|
||||
print(f"\n=== {base_tag} ===")
|
||||
top1, top5 = run_eval(base_eval)
|
||||
print(f"{base_tag} top1={top1:.2f}% top5={top5:.2f}%")
|
||||
results.append([base_tag, top1, top5])
|
||||
del base_eval
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
calib_batches = ttptq.build_calibration_loader(
|
||||
processor, args.calib_samples, args.batch_size, device, dtype
|
||||
)
|
||||
ttptq.quantize_with_recipe(model, args.recipe, calib_batches)
|
||||
|
||||
label = Path(args.recipe).stem # e.g. "fp8"
|
||||
tag = f"{label} ({runtime})"
|
||||
eval_model = to_eval_model(model, f"{label} model")
|
||||
print(f"\n=== {tag} ===")
|
||||
top1, top5 = run_eval(eval_model)
|
||||
print(f"{tag} top1={top1:.2f}% top5={top5:.2f}%")
|
||||
results.append([tag, top1, top5])
|
||||
|
||||
if args.results_path:
|
||||
with open(args.results_path, "w", newline="") as f:
|
||||
csv.writer(f).writerows(results)
|
||||
print(f"\nWrote results to {args.results_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,287 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""Quantize a HuggingFace ViT model with ModelOpt and compile with Torch-TensorRT.
|
||||
|
||||
Pipeline:
|
||||
|
||||
1. Load ``google/vit-large-patch16-224`` (`ViTForImageClassification`) from HF.
|
||||
2. Build a calibration loader from `zh-plus/tiny-imagenet` so the recipe runs
|
||||
end-to-end without ImageNet access.
|
||||
3. Run ``mtq.quantize`` with the ViT-specific FP8 recipe under
|
||||
`modelopt_recipes/huggingface/vit/ptq/`.
|
||||
4. Compile the quantized model with ``torch_tensorrt.compile(ir="dynamo",
|
||||
min_block_size=1)`` and verify the compiled-model argmax matches the
|
||||
fake-quant argmax on a sample input.
|
||||
|
||||
The quantized graph keeps Q/DQ nodes; the TRT compile step is what turns
|
||||
them into TRT precision layers.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from datasets import load_dataset
|
||||
from transformers import AutoImageProcessor, ViTForImageClassification
|
||||
|
||||
import modelopt.torch.opt as mto
|
||||
import modelopt.torch.quantization as mtq
|
||||
from modelopt.recipe import ModelOptPTQRecipe, load_recipe
|
||||
from modelopt.torch.quantization.utils import export_torch_mode
|
||||
|
||||
# Default ViT PTQ recipe under `modelopt_recipes/huggingface/vit/ptq/`. The
|
||||
# recipe loader resolves this relative path against the built-in recipe library;
|
||||
# pass `--recipe` for a different one.
|
||||
DEFAULT_RECIPE = "huggingface/vit/ptq/fp8"
|
||||
|
||||
|
||||
def load_model_and_processor(model_id: str, device: torch.device, dtype: torch.dtype):
|
||||
"""Pull the HF ViT classifier and its preprocessor."""
|
||||
print(f"Loading {model_id} (dtype={dtype})...")
|
||||
processor = AutoImageProcessor.from_pretrained(model_id)
|
||||
# `gelu_fast` selects the tanh-approximation GELU rather than the erf-based
|
||||
# default. Eager attention runs softmax through `F.softmax` instead of the
|
||||
# fused SDPA kernel, so the recipe's attention softmax-P quantizer
|
||||
# (`p_bmm_quantizer` on HF attention) is exercised during calibration and
|
||||
# emits Q/DQ around the softmax output on export.
|
||||
model = ViTForImageClassification.from_pretrained(
|
||||
model_id,
|
||||
torch_dtype=dtype,
|
||||
hidden_act="gelu_fast",
|
||||
attn_implementation="eager",
|
||||
)
|
||||
model.eval().to(device)
|
||||
return model, processor
|
||||
|
||||
|
||||
def build_calibration_loader(
|
||||
processor,
|
||||
num_samples: int,
|
||||
batch_size: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
"""Build a calibration tensor stream from tiny-imagenet."""
|
||||
print(f"Loading calibration data ({num_samples} samples)...")
|
||||
dataset = load_dataset("zh-plus/tiny-imagenet", split="train")
|
||||
dataset = dataset.shuffle(seed=42).select(range(num_samples))
|
||||
|
||||
tensors: list[torch.Tensor] = []
|
||||
for sample in dataset:
|
||||
image = sample["image"]
|
||||
if image.mode != "RGB":
|
||||
image = image.convert("RGB")
|
||||
pixel_values = processor(images=image, return_tensors="pt")["pixel_values"]
|
||||
tensors.append(pixel_values.squeeze(0))
|
||||
|
||||
batched = torch.stack(tensors).to(device=device, dtype=dtype)
|
||||
return torch.split(batched, batch_size)
|
||||
|
||||
|
||||
def quantize_with_recipe(model, recipe_path: str, calib_batches):
|
||||
"""Resolve the YAML recipe and run `mtq.quantize`."""
|
||||
print(f"Loading recipe: {recipe_path}")
|
||||
recipe = load_recipe(recipe_path)
|
||||
if not isinstance(recipe, ModelOptPTQRecipe):
|
||||
raise TypeError(f"Expected PTQ recipe, got {type(recipe).__name__}")
|
||||
quant_cfg = recipe.quantize.model_dump()
|
||||
|
||||
def forward_loop(model_):
|
||||
with torch.no_grad():
|
||||
for batch in calib_batches:
|
||||
model_(pixel_values=batch)
|
||||
|
||||
print("Running mtq.quantize ...")
|
||||
mtq.quantize(model, quant_cfg, forward_loop=forward_loop)
|
||||
mtq.print_quant_summary(model)
|
||||
return model
|
||||
|
||||
|
||||
class ViTLogitsWrapper(torch.nn.Module):
|
||||
"""Returns raw logits as a single tensor.
|
||||
|
||||
HF's `ViTForImageClassification.forward` returns an `ImageClassifierOutput`
|
||||
dataclass. `torch_tensorrt.compile` (and `torch.export`) need a tensor-tree
|
||||
return, so we unwrap it here.
|
||||
"""
|
||||
|
||||
def __init__(self, vit_model: torch.nn.Module):
|
||||
super().__init__()
|
||||
self.vit = vit_model
|
||||
|
||||
def forward(self, pixel_values: torch.Tensor) -> torch.Tensor:
|
||||
return self.vit(pixel_values=pixel_values).logits
|
||||
|
||||
|
||||
def compile_with_torch_tensorrt(model: torch.nn.Module, example_input: torch.Tensor):
|
||||
"""Compile the quantized model with Torch-TensorRT (Dynamo IR, strongly-typed)."""
|
||||
# Imported here (not at module scope) so the quantize-only `--skip_trt` path
|
||||
# still runs on hosts without torch_tensorrt installed.
|
||||
import torch_tensorrt
|
||||
|
||||
print("Compiling with torch_tensorrt.compile (Dynamo IR, dynamic batch)...")
|
||||
n, c, h, w = example_input.shape
|
||||
# torch.export specializes a size-1 dynamic dim to a constant, so trace at
|
||||
# opt batch >= 2; min=1 still serves batch 1 at runtime.
|
||||
opt_n = max(int(n), 2)
|
||||
with export_torch_mode(), torch_tensorrt.dynamo.Debugger(log_level="error"):
|
||||
trt_model = torch_tensorrt.compile(
|
||||
model,
|
||||
ir="dynamo",
|
||||
min_block_size=1,
|
||||
truncate_double=True,
|
||||
inputs=[
|
||||
torch_tensorrt.Input(
|
||||
min_shape=(1, c, h, w),
|
||||
opt_shape=(opt_n, c, h, w),
|
||||
max_shape=(1024, c, h, w),
|
||||
dtype=example_input.dtype,
|
||||
)
|
||||
],
|
||||
)
|
||||
return trt_model
|
||||
|
||||
|
||||
def dump_trt_layer_info(trt_model: torch.nn.Module, path: Path) -> None:
|
||||
"""Write the per-layer engine info of every TRT submodule to ``path``.
|
||||
|
||||
A Dynamo-compiled module can hold several ``TorchTensorRTModule`` subgraphs
|
||||
(the parts that fell back to PyTorch sit between them), so we concatenate the
|
||||
``get_layer_info()`` JSON of each.
|
||||
"""
|
||||
import torch_tensorrt
|
||||
|
||||
infos = [
|
||||
mod.get_layer_info()
|
||||
for _, mod in trt_model.named_modules()
|
||||
if isinstance(mod, torch_tensorrt.dynamo.runtime.TorchTensorRTModule)
|
||||
]
|
||||
if not infos:
|
||||
print("No TorchTensorRTModule found; nothing to dump (whole graph fell back to PyTorch?).")
|
||||
return
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text("\n".join(infos))
|
||||
print(f"Wrote TRT layer info ({len(infos)} engine(s)) to {path}")
|
||||
|
||||
|
||||
def _argmax_logits(out) -> torch.Tensor:
|
||||
"""Handle either an HF `ImageClassifierOutput` or a raw tensor."""
|
||||
logits = out.logits if hasattr(out, "logits") else out
|
||||
return logits.argmax(dim=-1)
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument(
|
||||
"--model_id",
|
||||
default="google/vit-large-patch16-224",
|
||||
help="HuggingFace model id of the ViT classifier to quantize.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--recipe",
|
||||
default=DEFAULT_RECIPE,
|
||||
help="Recipe path (relative to modelopt_recipes/ or an absolute YAML). "
|
||||
"Defaults to the ViT FP8 recipe.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--calib_samples",
|
||||
type=int,
|
||||
default=1024,
|
||||
help="Number of tiny-imagenet samples to use for calibration.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size",
|
||||
type=int,
|
||||
default=128,
|
||||
help="Batch size for calibration / TRT compile.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save_dir",
|
||||
type=str,
|
||||
default="./modelopt_quantized",
|
||||
help="Directory to save the quantized modelopt state-dict (FP16 weights "
|
||||
"+ Q/DQ metadata) — re-usable across runs without recalibration.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--skip_trt",
|
||||
action="store_true",
|
||||
help="Quantize + run the fake-quant model only; skip torch_tensorrt.compile. "
|
||||
"Useful for environments without torch_tensorrt installed.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--layer_info_path",
|
||||
default=None,
|
||||
help="If set, write the compiled TRT engine's per-layer info "
|
||||
"(get_layer_info()) to this file.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
raise SystemExit("This example requires a CUDA-capable GPU.")
|
||||
device = torch.device("cuda")
|
||||
dtype = torch.float16
|
||||
|
||||
model, processor = load_model_and_processor(args.model_id, device, dtype)
|
||||
image_size = model.config.image_size
|
||||
num_channels = model.config.num_channels
|
||||
example_input = torch.randn(
|
||||
args.batch_size, num_channels, image_size, image_size, device=device, dtype=dtype
|
||||
)
|
||||
|
||||
print("\n=== Baseline (FP16) ===")
|
||||
with torch.no_grad():
|
||||
baseline_pred = _argmax_logits(model(example_input))
|
||||
print(f"Baseline argmax class: {baseline_pred.tolist()}")
|
||||
|
||||
calib_batches = build_calibration_loader(
|
||||
processor, args.calib_samples, args.batch_size, device, dtype
|
||||
)
|
||||
|
||||
quantize_with_recipe(model, args.recipe, calib_batches)
|
||||
|
||||
save_path = Path(args.save_dir)
|
||||
save_path.mkdir(parents=True, exist_ok=True)
|
||||
ckpt = save_path / "vit_modelopt_state.pt"
|
||||
mto.save(model, ckpt)
|
||||
print(f"Saved quantized modelopt state to {ckpt}")
|
||||
|
||||
print("\n=== Fake-quant (modelopt) ===")
|
||||
with torch.no_grad():
|
||||
fq_pred = _argmax_logits(model(example_input))
|
||||
fq_match = (fq_pred == baseline_pred).all().item()
|
||||
print(f"Quantized argmax class: {fq_pred.tolist()} (matches baseline: {fq_match})")
|
||||
|
||||
if args.skip_trt:
|
||||
print("\n--skip_trt set; not compiling with Torch-TensorRT.")
|
||||
return
|
||||
|
||||
wrapped = ViTLogitsWrapper(model).to(device).eval()
|
||||
trt_model = compile_with_torch_tensorrt(wrapped, example_input)
|
||||
|
||||
if args.layer_info_path:
|
||||
dump_trt_layer_info(trt_model, Path(args.layer_info_path))
|
||||
|
||||
print("\n=== Torch-TensorRT compiled ===")
|
||||
with torch.no_grad():
|
||||
trt_pred = trt_model(example_input).argmax(dim=-1)
|
||||
trt_match = (trt_pred == baseline_pred).all().item()
|
||||
print(f"TRT argmax class: {trt_pred.tolist()} (matches baseline: {trt_match})")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user