UH-OH

It looks like you don’t have access to that feature yet

Contact sales to get upgraded to the full DevStudio experience.

UH-OH

It looks like you don't have access to that feature yet.

Introduction to the Chimera SDK
Chimera SDK Quick Start Guide
Chimera SDK Command Line Interface (CLI)
Tutorial: Using SDK as a Library
Tutorials & Model Demos
Quantization Tutorials
Multicore Demo
Chimera LLVM C++ Compiler
Chimera SDK Licensing Policy Documentation
Glossary
Chimera Software User GuideTutorials & Model DemosQuantization TutorialsQAT: EfficientNet Quantization Aware Training

QAT: EfficientNet Quantization Aware Training


NOTE: The Jupyter Notebook below is included in the Chimera SDK and can be run interactively by running the following CLI command:

$ quadric sdk notebook

From the Jupyter Notebook window in your browser, select the notebook named /quadric/sdk-cli/examples/quantization/QAT.ipynb.


PyTorch QAT -> ONNX Runtime Pipeline

In this tutorial we allow 2 paths: (1) use the PyTorch 2 Export (pt2e) library to perform quantization-aware training (QAT) on EfficientNet-B7, and export it such that it can be run through ONNX Runtime and (2) export a pre-trained QAT model from PyTorch so that it can be lowered in CGC.

Notes:

  • For path (1), we will use torchvision.models.efficientnet_b7 and for path (2) we will use torchvision.models.quantization.mobilenet_v2. We provide a flag to toggle between these in the Imports and Setup section.
  • Training is expected to be done with a GPU, which can be mounted to the docker container with the option --gpus all or --gpus device=0 (replace 0 with whichever device you'd like to use). If training is not feasible due to compute, time, etc... constraints, we have provided the pre-trained model in examples/models/efficientnet/efficientnet-quadric-qat.onnx which can be validated using step (8)
  • pt2e is still in prototype phase (as of 10/08/2024), breaking changes may occur later
  • Other PyTorch quantization libraries are available, but we currently make no guarantee the process will work as intended using them
  • Each network requires slightly different post-processing, so not all networks may be supported yet with the post-processing steps currently implemented
  • Images, labels, and models used in this notebook conform to the ImageNet-1K standard (i.e. use 224x224 image resolution) and have subjects represented in the 1000 ImageNet classes; any different input size or class labels will require user-supplied datasets

High Level Overview

  1. Set up training helper functions
  2. Define our quantization configuration
  3. Prepare the training dataset
  4. Load the model from the torchvision library
  5. Perform QAT using the pt2e library
  6. ONNX export and post-processing
  7. Prepare the validation dataset
  8. Validate and compare against FP32 model
  9. CGC lowering (pretrained MobileNet only)

0. Imports and setup

At the bottom of the cell, we include a flag to choose whether to continue training EfficientNet or to skip training and use a pre-trained MobileNet.

%pip install -r ../requirements_gpu.txt
Requirement already satisfied: torch==2.5.0 in /usr/local/lib/python3.10/dist-packages (from -r ../requirements_gpu.txt (line 1)) (2.5.0)
Requirement already satisfied: torchvision==0.20.0 in /usr/local/lib/python3.10/dist-packages (from -r ../requirements_gpu.txt (line 2)) (0.20.0)
Requirement already satisfied: s3fs in /usr/local/lib/python3.10/dist-packages (from -r ../requirements_gpu.txt (line 3)) (2026.6.0)
Requirement already satisfied: onnxscript==0.2.0 in /usr/local/lib/python3.10/dist-packages (from -r ../requirements_gpu.txt (line 4)) (0.2.0)
Requirement already satisfied: nvidia-cusolver-cu12==11.6.1.9 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (11.6.1.9)
Requirement already satisfied: nvidia-nvjitlink-cu12==12.4.127 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (12.4.127)
Requirement already satisfied: nvidia-cuda-runtime-cu12==12.4.127 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (12.4.127)
Requirement already satisfied: jinja2 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (3.1.6)
Requirement already satisfied: nvidia-cufft-cu12==11.2.1.3 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (11.2.1.3)
Requirement already satisfied: triton==3.1.0 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (3.1.0)
Requirement already satisfied: nvidia-curand-cu12==10.3.5.147 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (10.3.5.147)
Requirement already satisfied: sympy==1.13.1 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (1.13.1)
Requirement already satisfied: typing-extensions>=4.8.0 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (4.16.0)
Requirement already satisfied: filelock in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (3.16.1)
Requirement already satisfied: fsspec in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (2026.6.0)
Requirement already satisfied: nvidia-cusparse-cu12==12.3.1.170 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (12.3.1.170)
Requirement already satisfied: nvidia-cuda-nvrtc-cu12==12.4.127 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (12.4.127)
Requirement already satisfied: nvidia-nccl-cu12==2.21.5 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (2.21.5)
Requirement already satisfied: nvidia-nvtx-cu12==12.4.127 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (12.4.127)
Requirement already satisfied: nvidia-cuda-cupti-cu12==12.4.127 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (12.4.127)
Requirement already satisfied: nvidia-cudnn-cu12==9.1.0.70 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (9.1.0.70)
Requirement already satisfied: networkx in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (2.8.5)
Requirement already satisfied: nvidia-cublas-cu12==12.4.5.8 in /usr/local/lib/python3.10/dist-packages (from torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (12.4.5.8)
Requirement already satisfied: pillow!=8.3.*,>=5.3.0 in /usr/local/lib/python3.10/dist-packages (from torchvision==0.20.0->-r ../requirements_gpu.txt (line 2)) (12.3.0)
Requirement already satisfied: numpy in /usr/local/lib/python3.10/dist-packages (from torchvision==0.20.0->-r ../requirements_gpu.txt (line 2)) (1.24.4)
Requirement already satisfied: ml_dtypes in /usr/local/lib/python3.10/dist-packages (from onnxscript==0.2.0->-r ../requirements_gpu.txt (line 4)) (0.3.2)
Requirement already satisfied: onnx>=1.16 in /usr/local/lib/python3.10/dist-packages (from onnxscript==0.2.0->-r ../requirements_gpu.txt (line 4)) (1.16.2)
Requirement already satisfied: packaging in /usr/local/lib/python3.10/dist-packages (from onnxscript==0.2.0->-r ../requirements_gpu.txt (line 4)) (26.2)
Requirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.10/dist-packages (from sympy==1.13.1->torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (1.3.0)
Requirement already satisfied: aiobotocore<4.0.0,>=2.19.0 in /usr/local/lib/python3.10/dist-packages (from s3fs->-r ../requirements_gpu.txt (line 3)) (3.8.0)
Requirement already satisfied: aiohttp!=4.0.0a0,!=4.0.0a1,>=3.9.0 in /usr/local/lib/python3.10/dist-packages (from s3fs->-r ../requirements_gpu.txt (line 3)) (3.14.1)
Requirement already satisfied: multidict<7.0.0,>=6.0.0 in /usr/local/lib/python3.10/dist-packages (from aiobotocore<4.0.0,>=2.19.0->s3fs->-r ../requirements_gpu.txt (line 3)) (6.7.1)
Requirement already satisfied: python-dateutil<3.0.0,>=2.1 in /usr/local/lib/python3.10/dist-packages (from aiobotocore<4.0.0,>=2.19.0->s3fs->-r ../requirements_gpu.txt (line 3)) (2.9.0.post0)
Requirement already satisfied: wrapt<3.0.0,>=1.10.10 in /usr/local/lib/python3.10/dist-packages (from aiobotocore<4.0.0,>=2.19.0->s3fs->-r ../requirements_gpu.txt (line 3)) (2.2.2)
Requirement already satisfied: botocore<1.43.47,>=1.43.3 in /usr/local/lib/python3.10/dist-packages (from aiobotocore<4.0.0,>=2.19.0->s3fs->-r ../requirements_gpu.txt (line 3)) (1.43.46)
Requirement already satisfied: aioitertools<1.0.0,>=0.5.1 in /usr/local/lib/python3.10/dist-packages (from aiobotocore<4.0.0,>=2.19.0->s3fs->-r ../requirements_gpu.txt (line 3)) (0.13.0)
Requirement already satisfied: jmespath<2.0.0,>=0.7.1 in /usr/local/lib/python3.10/dist-packages (from aiobotocore<4.0.0,>=2.19.0->s3fs->-r ../requirements_gpu.txt (line 3)) (1.1.0)
Requirement already satisfied: attrs>=17.3.0 in /usr/local/lib/python3.10/dist-packages (from aiohttp!=4.0.0a0,!=4.0.0a1,>=3.9.0->s3fs->-r ../requirements_gpu.txt (line 3)) (26.1.0)
Requirement already satisfied: aiosignal>=1.4.0 in /usr/local/lib/python3.10/dist-packages (from aiohttp!=4.0.0a0,!=4.0.0a1,>=3.9.0->s3fs->-r ../requirements_gpu.txt (line 3)) (1.4.0)
Requirement already satisfied: async-timeout<6.0,>=4.0 in /usr/local/lib/python3.10/dist-packages (from aiohttp!=4.0.0a0,!=4.0.0a1,>=3.9.0->s3fs->-r ../requirements_gpu.txt (line 3)) (5.0.1)
Requirement already satisfied: frozenlist>=1.1.1 in /usr/local/lib/python3.10/dist-packages (from aiohttp!=4.0.0a0,!=4.0.0a1,>=3.9.0->s3fs->-r ../requirements_gpu.txt (line 3)) (1.8.0)
Requirement already satisfied: propcache>=0.2.0 in /usr/local/lib/python3.10/dist-packages (from aiohttp!=4.0.0a0,!=4.0.0a1,>=3.9.0->s3fs->-r ../requirements_gpu.txt (line 3)) (0.5.2)
Requirement already satisfied: yarl<2.0,>=1.17.0 in /usr/local/lib/python3.10/dist-packages (from aiohttp!=4.0.0a0,!=4.0.0a1,>=3.9.0->s3fs->-r ../requirements_gpu.txt (line 3)) (1.24.2)
Requirement already satisfied: aiohappyeyeballs>=2.5.0 in /usr/local/lib/python3.10/dist-packages (from aiohttp!=4.0.0a0,!=4.0.0a1,>=3.9.0->s3fs->-r ../requirements_gpu.txt (line 3)) (2.7.1)
Requirement already satisfied: protobuf>=3.20.2 in /usr/local/lib/python3.10/dist-packages (from onnx>=1.16->onnxscript==0.2.0->-r ../requirements_gpu.txt (line 4)) (4.25.3)
Requirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.10/dist-packages (from jinja2->torch==2.5.0->-r ../requirements_gpu.txt (line 1)) (3.0.3)
Requirement already satisfied: urllib3!=2.2.0,<3,>=1.25.4 in /usr/local/lib/python3.10/dist-packages (from botocore<1.43.47,>=1.43.3->aiobotocore<4.0.0,>=2.19.0->s3fs->-r ../requirements_gpu.txt (line 3)) (1.26.20)
Requirement already satisfied: six>=1.5 in /usr/local/lib/python3.10/dist-packages (from python-dateutil<3.0.0,>=2.1->aiobotocore<4.0.0,>=2.19.0->s3fs->-r ../requirements_gpu.txt (line 3)) (1.17.0)
Requirement already satisfied: idna>=2.0 in /usr/local/lib/python3.10/dist-packages (from yarl<2.0,>=1.17.0->aiohttp!=4.0.0a0,!=4.0.0a1,>=3.9.0->s3fs->-r ../requirements_gpu.txt (line 3)) (3.18)
WARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv
Note: you may need to restart the kernel to use updated packages.
import os
import sys
import time
import copy
import logging
import itertools
import warnings
import tempfile
import random
from pathlib import Path
from dataclasses import dataclass
from datasets import load_from_disk
from typing import Optional
import numpy as np
from PIL import Image

import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.onnx
from torch.utils.data import DataLoader, Subset
from torch._export import capture_pre_autograd_graph
from torch.export import export
from torch.ao.quantization.quantize_pt2e import (
    prepare_qat_pt2e,
    convert_pt2e,
)
from torch.ao.quantization.quantizer import (
    Quantizer,
    QuantizationSpec,
    FixedQParamsQuantizationSpec,
    QuantizationAnnotation,
)
from torch.ao.quantization.fake_quantize import FakeQuantize
from torch.ao.quantization.quantizer import QuantizationSpec
from torch.ao.quantization.observer import MinMaxObserver
from torch.fx.passes.utils.source_matcher_utils import get_source_partitions

import torchvision
from torchvision.datasets import ImageFolder
from torchvision.models import (
    efficientnet_b7,
    EfficientNet_B7_Weights,
    MobileNet_V2_Weights,
)
from torchvision.models.quantization import (
    mobilenet_v2,
    MobileNet_V2_QuantizedWeights,
)
from torchvision.transforms import (
    RandomResizedCrop,
    RandomHorizontalFlip,
    CenterCrop,
    Compose,
    Normalize,
    Resize,
    ToTensor,
)

import sdk_cli.lib.qat_processor as processor
from tvm.contrib.epu.chimera_job.chimera_job import ChimeraJob
from tvm.contrib.epu.chimera_job.hw_config import HWConfig
from tvm.contrib.epu.chimera_job.constants import DEFAULT_ONNX_OPSET
from sdk_cli.utils.dataloaders import CalibrationDataLoader
from sdk_cli.utils.datasets import ImageNet_Mini_Quadric, QuadricCalibration
from sdk_cli.utils.datasets.ImageNet import (
    IMAGENET_1K_NORMALIZATION_PARAMETERS,
    OPTIMIZED_IMAGENET_1K_LABELS,
)
from sdk_cli.utils.model_helpers import ClassifyResult
from sdk_cli.utils.performance_trackers import ClassifierPerformanceTracker
from sdk_cli.lib.quantize import (
    QuantizationExperiment,
    QuantizedONNXModel,
    run_quantization_experiment,
)

warnings.filterwarnings(action="ignore", category=DeprecationWarning, module=r".*")
warnings.filterwarnings(action="default", module=r"torch.ao.quantization")

## NOTE: below seeds manually set for reproducability
torch.manual_seed(191009)
np.random.seed(2147483648)
random.seed(2147483648)
torch.use_deterministic_algorithms(True)
os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"

pretrained_mobilenet = False
if not pretrained_mobilenet:
    torch.set_default_device("cuda")
else:
    torch.set_default_device("cpu")

1. Set up training helper functions

These will assist us in creating our training loop. The pt2e documentation page listed above also gives ways to perform evaluation mid-training, which we have omitted here as we will be evaluating the model in ONNX Runtime.

class AverageMeter(object):
    def __init__(self, name, fmt=":f"):
        self.name = name
        self.fmt = fmt
        self.reset()

    def reset(self):
        self.val = 0
        self.avg = 0
        self.sum = 0
        self.count = 0

    def update(self, val, n=1):
        self.val = val
        self.sum += val * n
        self.count += n
        self.avg = self.sum / self.count

    def __str__(self):
        return f"{name} {val:f} ({avg:f})"


def accuracy(output, target, topk=(1,)):
    with torch.no_grad():
        maxk = max(topk)
        batch_size = target.size(0)

        _, pred = output.topk(maxk, 1, True, True)
        pred = pred.t()
        correct = pred.eq(target.view(1, -1).expand_as(pred))

        res = []
        for k in topk:
            correct_k = correct[:k].reshape(-1).float().sum(0, keepdim=True)
            res.append(correct_k.mul_(100.0 / batch_size))
        return res


def train_one_epoch(model, criterion, optimizer, data_loader, device, ntrain_batches):
    top1 = AverageMeter("Acc@1")
    top5 = AverageMeter("Acc@5")
    avgloss = AverageMeter("Loss")

    cnt = 0
    for batch in data_loader:
        start_time = time.time()
        image, target = batch.values()
        print(".", end="")
        cnt += 1
        image = image.to(device)
        target = target.to(device)
        output = model(image)
        loss = criterion(output, target)
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()
        acc1, acc5 = accuracy(output, target, topk=(1, 5))
        top1.update(acc1[0], image.size(0))
        top5.update(acc5[0], image.size(0))
        avgloss.update(loss, image.size(0))
        if cnt >= ntrain_batches:
            print("Loss", avgloss.avg)
            print(
                "Training: * Acc@1 {top1.avg:.3f} Acc@5 {top5.avg:.3f}".format(top1=top1, top5=top5)
            )
            return

    print(
        "Full imagenet train set:  * Acc@1 {top1.global_avg:.3f} Acc@5 {top5.global_avg:.3f}".format(
            top1=top1, top5=top5
        )
    )


def evaluate(model, data_loader, device, neval_batches):
    top1 = AverageMeter("Acc@1", ":6.2f")
    top5 = AverageMeter("Acc@5", ":6.2f")
    cnt = 0
    with torch.no_grad():
        for image, target in data_loader:
            image = image.to(device)
            target = target.to(device)
            output = model(image)
            cnt += 1
            acc1, acc5 = accuracy(output, target, topk=(1, 5))
            top1.update(acc1[0], image.size(0))
            top5.update(acc5[0], image.size(0))
            if cnt >= neval_batches:
                return top1, top5
    print("")

    return top1, top5

2. Define our quantization configuration

The pt2e quantization library is extremely flexible. It allows us to define the method of quantization for each tensor of each individual operation. Here we will define a simple quantization configuration, but you may choose to experiment with the options available.

@dataclass(eq=True, frozen=True)
class QuantizationConfig:
    input_activation: Optional[QuantizationSpec]
    output_activation: Optional[QuantizationSpec]
    weight: Optional[QuantizationSpec]
    bias: Optional[QuantizationSpec]
    is_qat: bool = False


quadric_weight_fake_quant = FakeQuantize.with_args(
    observer=MinMaxObserver,
    quant_min=-128,
    quant_max=127,
    dtype=torch.qint8,
    qscheme=torch.per_tensor_symmetric,
)

quadric_activation_fake_quant = FakeQuantize.with_args(
    observer=MinMaxObserver,
    quant_min=-128,
    quant_max=127,
    dtype=torch.qint8,
    qscheme=torch.per_tensor_affine,
)


class QuadricQuantizer(Quantizer):
    def __init__(self):
        super().__init__()
        self.global_config: QuantizationConfig = None

    def set_global_config(self, quant_config: QuantizationConfig):
        self.global_config = quant_config
        return self

    def get_default_config(self):
        activation_spec = QuantizationSpec(
            observer_or_fake_quant_ctr=quadric_activation_fake_quant,
            dtype=torch.int8,
            quant_min=-128,
            quant_max=127,
            qscheme=torch.per_tensor_affine,
        )
        weight_spec = QuantizationSpec(
            observer_or_fake_quant_ctr=quadric_weight_fake_quant,
            dtype=torch.int8,
            quant_min=-128,
            quant_max=127,
            qscheme=torch.per_tensor_symmetric,  # NOTE: currently only supporting symmetric weight quantization
        )
        quant_config = QuantizationConfig(
            input_activation=activation_spec,
            output_activation=activation_spec,
            weight=weight_spec,
            bias=None,
            is_qat=True,
        )
        return quant_config

    def annotate(self, model: torch.fx.GraphModule) -> torch.fx.GraphModule:
        self._annotate_conv2d(model)
        self._annotate_gemm(model)
        self._annotate_average_pool(model)
        self._annotate_batch_norm(model)
        return model

    def _annotate_conv2d(self, model: torch.fx.GraphModule):
        conv_partitions = get_source_partitions(model.graph, [nn.Conv2d, F.conv2d])
        conv_partitions = list(itertools.chain(*conv_partitions.values()))
        for partition in conv_partitions:
            conv_node = partition.output_nodes[0]
            input_qspec_map = {
                conv_node.args[0]: self.global_config.input_activation,
                conv_node.args[1]: self.global_config.weight,
            }
            conv_node.meta["quantization_annotation"] = QuantizationAnnotation(
                input_qspec_map=input_qspec_map,
                output_qspec=self.global_config.output_activation,
                _annotated=True,
            )

    def _annotate_gemm(self, model: torch.fx.GraphModule):
        gemm_partitions = get_source_partitions(
            model.graph, [torch.matmul, torch.mm, nn.Linear, F.linear]
        )
        gemm_partitions = list(itertools.chain(*gemm_partitions.values()))
        for partition in gemm_partitions:
            gemm_node = partition.output_nodes[0]
            input_qspec_map = {
                gemm_node.args[0]: self.global_config.input_activation,
                gemm_node.args[1]: self.global_config.input_activation,
            }
            gemm_node.meta["quantization_annotation"] = QuantizationAnnotation(
                input_qspec_map=input_qspec_map,
                output_qspec=self.global_config.output_activation,
                _annotated=True,
            )

    def _annotate_average_pool(self, model: torch.fx.GraphModule):
        gap_partitions = get_source_partitions(
            model.graph,
            [nn.AvgPool2d, nn.AdaptiveAvgPool2d, F.avg_pool2d, F.adaptive_avg_pool2d],
        )
        gap_partitions = list(itertools.chain(*gap_partitions.values()))
        for partition in gap_partitions:
            gap_node = partition.output_nodes[0]
            input_qspec_map = {
                gap_node.args[0]: self.global_config.input_activation,
            }
            gap_node.meta["quantization_annotation"] = QuantizationAnnotation(
                input_qspec_map=input_qspec_map,
                output_qspec=self.global_config.output_activation,
                _annotated=True,
            )

    def _annotate_batch_norm(self, model: torch.fx.GraphModule):
        bn_partitions = get_source_partitions(model.graph, [nn.BatchNorm2d, F.batch_norm])
        bn_partitions = list(itertools.chain(*bn_partitions.values()))
        for partition in bn_partitions:
            bn_node = partition.output_nodes[0]
            input_qspec_map = {
                bn_node.args[0]: self.global_config.input_activation,
            }
            bn_node.meta["quantization_annotation"] = QuantizationAnnotation(
                input_qspec_map=input_qspec_map,
                output_qspec=self.global_config.output_activation,
            )

    def validate(self, model):
        pass

    @classmethod
    def get_supported_operators(cls):
        return []

3. Prepare the training dataset

We use images from ImageNet-1K to perform the training and validation. All input tensors in this dataset are of dimension NCHW format [3, 224, 224] and need to be float-normalized with a mean, sigma of [0.485, 0.456, 0.406], [0.229, 0.224, 0.225]. For training, we include the transformations RandomResizedCrop and RandomHorizontalFlip for robustness.

normalize = Normalize(
    IMAGENET_1K_NORMALIZATION_PARAMETERS.channel_means,
    IMAGENET_1K_NORMALIZATION_PARAMETERS.channel_standard_deviations,
)
train_transforms = Compose(
    [
        RandomResizedCrop(224),
        RandomHorizontalFlip(),
        ToTensor(),
        normalize,
    ]
)
test_transforms = Compose(
    [
        Resize(224),
        CenterCrop(224),
        ToTensor(),
        normalize,
    ]
)


def apply_transforms(imgs):
    imgs["image"] = [train_transforms(img.convert("RGB")) for img in imgs["image"]]
    return imgs


if not pretrained_mobilenet:
    train_set = load_from_disk(
        "s3://sdk-cli-datasets/imagenet_1k_val_train.hf/validation",
        storage_options={"anon": True},
    )
    train_set.set_transform(apply_transforms)
    test_set = ImageFolder(
        root="../common/validation/imagenet-mini-quadric/val", transform=test_transforms
    )

    train_batch_size = 8
    test_batch_size = 8
    generator_train = torch.Generator(device="cuda")
    generator_train.manual_seed(67280421310721)
    generator_test = torch.Generator(device="cuda")
    generator_test.manual_seed(672804213107421)
    sampler_train = torch.utils.data.RandomSampler(train_set, generator=generator_train)
    sampler_test = torch.utils.data.SequentialSampler(test_set)
    data_loader_train = torch.utils.data.DataLoader(
        train_set,
        batch_size=train_batch_size,
        generator=generator_train,
        sampler=sampler_train,
    )
    data_loader_test = torch.utils.data.DataLoader(
        test_set,
        batch_size=test_batch_size,
        generator=generator_test,
        sampler=sampler_test,
    )
/usr/local/lib/python3.10/dist-packages/datasets/table.py:1421: FutureWarning: promote has been superseded by promote_options='default'.
  table = cls._concat_blocks(blocks, axis=0)

4. Load the model from torchvision

In this tutorial we are using EfficientNet-B7, but other networks from the torchvision library may be used as well. The original graph goes through a 2-step process using capture_pre_autograd_graph and the previously defined QuadricQuantizer in order to be ready for QAT.

For more documentation about the models supported please see: https://pytorch.org/vision/0.8/models.html

if not pretrained_mobilenet:
    float_model = efficientnet_b7(weights=EfficientNet_B7_Weights.DEFAULT)

    example_input = torch.rand(1, 3, 224, 224)
    exported_model = capture_pre_autograd_graph(float_model, (example_input,))

    quantizer = QuadricQuantizer()
    quantizer.set_global_config(quantizer.get_default_config())
    prepared_model = prepare_qat_pt2e(exported_model, quantizer)
    for n in prepared_model.graph.nodes:
        if n.target == torch.ops.aten._native_batch_norm_legit.default:
            n.target = torch.ops.aten.cudnn_batch_norm.default
    _ = prepared_model.recompile()
else:
    float_model = mobilenet_v2(
        weights=MobileNet_V2_QuantizedWeights.IMAGENET1K_QNNPACK_V1, quantize=True
    )
    example_input = torch.rand(1, 3, 224, 224)
W0718 12:40:28.872000 101 torch/_export/__init__.py:64] +============================+
W0718 12:40:28.873000 101 torch/_export/__init__.py:65] |     !!!   WARNING   !!!    |
W0718 12:40:28.873000 101 torch/_export/__init__.py:66] +============================+
W0718 12:40:28.874000 101 torch/_export/__init__.py:67] capture_pre_autograd_graph() is deprecated and doesn't provide any function guarantee moving forward.
W0718 12:40:28.874000 101 torch/_export/__init__.py:68] Please switch to use torch.export.export_for_training instead.
/usr/local/lib/python3.10/dist-packages/onnxscript/converter.py:823: FutureWarning: 'onnxscript.values.Op.param_schemas' is deprecated in version 0.1 and will be removed in the future. Please use '.op_signature' instead.
  param_schemas = callee.param_schemas()
/usr/local/lib/python3.10/dist-packages/onnxscript/converter.py:823: FutureWarning: 'onnxscript.values.OnnxFunction.param_schemas' is deprecated in version 0.1 and will be removed in the future. Please use '.op_signature' instead.
  param_schemas = callee.param_schemas()

5. Perform QAT using the pt2e library

We follow the documentation given here. We start by defining some hyperparameters here. You can alter these parameters as you like, but a good starting point is these Nvidia docs.

if not pretrained_mobilenet:
    num_epochs = 275
    num_train_batches = 32
    num_eval_batches = 256
    num_observer_update_epochs = 1000
    num_batch_norm_update_epochs = 0
    num_epochs_between_evals = 50

    criterion = nn.CrossEntropyLoss()
    optimizer = torch.optim.AdamW(prepared_model.parameters(), lr=2e-6, weight_decay=1e-5)

Now, we move on to the training loop. We trained using one GPU, so the results may vary from what we observed when using more.

if not pretrained_mobilenet:
    num_observer_update_flag = True
    num_batch_norm_update_flag = True

    for epoch in range(num_epochs):
        if epoch + 1 >= num_observer_update_epochs and num_observer_update_flag:
            print("Disabling observer for subseq epochs, epoch = ", epoch)
            prepared_model.apply(torch.ao.quantization.disable_observer)
            num_observer_update_flag = False

        if epoch + 1 >= num_batch_norm_update_epochs and num_batch_norm_update_flag:
            print("Freezing BN for subseq epochs, epoch = ", epoch)
            for n in prepared_model.graph.nodes:
                if n.target in [
                    torch.ops.aten._native_batch_norm_legit.default,
                    torch.ops.aten.cudnn_batch_norm.default,
                ]:
                    new_args = list(n.args)
                    new_args[5] = False
                    n.args = tuple(new_args)
            prepared_model.recompile()
            num_batch_norm_update_flag = False

        train_one_epoch(
            prepared_model,
            criterion,
            optimizer,
            data_loader_train,
            "cuda",
            num_train_batches,
        )
        print("^ Epoch: %d" % (epoch + 1))

        if (epoch + 1) % num_epochs_between_evals == 0:
            prepared_model_copy = copy.deepcopy(prepared_model)
            torch.set_default_device("cpu")
            prepared_model_copy = prepared_model_copy.to("cpu")
            prepared_model_copy.recompile()

            quantized_model = convert_pt2e(prepared_model_copy)
            torch.set_default_device("cuda")
            quantized_model = quantized_model.to("cuda")
            quantized_model.recompile()
            torch.ao.quantization.move_exported_model_to_eval(quantized_model)

            top1, top5 = evaluate(
                quantized_model,
                data_loader_test,
                device="cuda",
                neval_batches=num_eval_batches,
            )
            print(
                "Epoch %d: Evaluation accuracy on %d images, %2.2f"
                % (epoch + 1, num_eval_batches * test_batch_size, top1.avg)
            )
Freezing BN for subseq epochs, epoch =  0
................................Loss tensor(2.1282, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 57.812 Acc@5 79.297
^ Epoch: 1
................................Loss tensor(2.2726, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 56.641 Acc@5 75.000
^ Epoch: 2
................................Loss tensor(2.4127, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 55.859 Acc@5 75.000
^ Epoch: 3
................................Loss tensor(2.3027, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 55.859 Acc@5 77.734
^ Epoch: 4
................................Loss tensor(2.4992, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 50.391 Acc@5 75.000
^ Epoch: 5
................................Loss tensor(2.2424, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.156 Acc@5 76.172
^ Epoch: 6
................................Loss tensor(2.4783, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 55.859 Acc@5 73.828
^ Epoch: 7
................................Loss tensor(2.2716, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 58.594 Acc@5 78.516
^ Epoch: 8
................................Loss tensor(2.3638, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 52.734 Acc@5 77.734
^ Epoch: 9
................................Loss tensor(2.0994, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 57.422 Acc@5 78.125
^ Epoch: 10
................................Loss tensor(2.1096, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 57.031 Acc@5 81.250
^ Epoch: 11
................................Loss tensor(1.9826, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.938 Acc@5 79.297
^ Epoch: 12
................................Loss tensor(2.2434, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 58.594 Acc@5 76.562
^ Epoch: 13
................................Loss tensor(2.0389, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 61.719 Acc@5 80.469
^ Epoch: 14
................................Loss tensor(2.0520, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 62.109 Acc@5 78.906
^ Epoch: 15
................................Loss tensor(2.1003, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 58.984 Acc@5 81.250
^ Epoch: 16
................................Loss tensor(2.0667, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.938 Acc@5 79.688
^ Epoch: 17
................................Loss tensor(2.2938, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 53.906 Acc@5 78.125
^ Epoch: 18
................................Loss tensor(2.0455, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 58.594 Acc@5 80.859
^ Epoch: 19
................................Loss tensor(2.1940, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 55.078 Acc@5 77.344
^ Epoch: 20
................................Loss tensor(1.9785, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 58.594 Acc@5 80.469
^ Epoch: 21
................................Loss tensor(2.0548, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 57.422 Acc@5 79.297
^ Epoch: 22
................................Loss tensor(2.0890, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.156 Acc@5 78.906
^ Epoch: 23
................................Loss tensor(1.9914, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 58.594 Acc@5 79.297
^ Epoch: 24
................................Loss tensor(2.0220, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.156 Acc@5 79.688
^ Epoch: 25
................................Loss tensor(1.8744, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 61.328 Acc@5 80.078
^ Epoch: 26
................................Loss tensor(1.6714, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.062 Acc@5 85.156
^ Epoch: 27
................................Loss tensor(1.7386, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 65.234 Acc@5 83.203
^ Epoch: 28
................................Loss tensor(2.0804, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 55.859 Acc@5 81.250
^ Epoch: 29
................................Loss tensor(1.9936, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 58.984 Acc@5 82.422
^ Epoch: 30
................................Loss tensor(1.8127, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.062 Acc@5 83.594
^ Epoch: 31
................................Loss tensor(2.0115, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 57.422 Acc@5 78.906
^ Epoch: 32
................................Loss tensor(1.9149, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.156 Acc@5 80.078
^ Epoch: 33
................................Loss tensor(1.8975, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.156 Acc@5 80.859
^ Epoch: 34
................................Loss tensor(1.6893, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.062 Acc@5 85.156
^ Epoch: 35
................................Loss tensor(1.5138, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.062 Acc@5 85.938
^ Epoch: 36
................................Loss tensor(2.0614, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.547 Acc@5 77.344
^ Epoch: 37
................................Loss tensor(1.8944, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.547 Acc@5 82.422
^ Epoch: 38
................................Loss tensor(1.6987, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 65.234 Acc@5 83.984
^ Epoch: 39
................................Loss tensor(1.8844, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 62.109 Acc@5 82.422
^ Epoch: 40
................................Loss tensor(1.7904, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 62.109 Acc@5 82.031
^ Epoch: 41
................................Loss tensor(1.8081, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 63.281 Acc@5 81.641
^ Epoch: 42
................................Loss tensor(1.6973, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 61.328 Acc@5 83.594
^ Epoch: 43
................................Loss tensor(1.8584, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 56.641 Acc@5 84.375
^ Epoch: 44
................................Loss tensor(1.8433, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 62.109 Acc@5 80.859
^ Epoch: 45
................................Loss tensor(1.7137, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.016 Acc@5 82.812
^ Epoch: 46
................................Loss tensor(1.7836, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.453 Acc@5 83.203
^ Epoch: 47
................................Loss tensor(1.5613, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.188 Acc@5 84.375
^ Epoch: 48
................................Loss tensor(1.6796, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 63.281 Acc@5 82.812
^ Epoch: 49
................................Loss tensor(1.9394, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 63.281 Acc@5 77.734
^ Epoch: 50
Epoch 50: Evaluation accuracy on 2048 images, 72.90
................................Loss tensor(1.9104, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 58.594 Acc@5 79.297
^ Epoch: 51
................................Loss tensor(1.8895, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 58.984 Acc@5 77.734
^ Epoch: 52
................................Loss tensor(1.4391, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.578 Acc@5 85.156
^ Epoch: 53
................................Loss tensor(1.6142, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.453 Acc@5 85.156
^ Epoch: 54
................................Loss tensor(1.6646, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.406 Acc@5 83.203
^ Epoch: 55
................................Loss tensor(1.7046, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 62.500 Acc@5 83.984
^ Epoch: 56
................................Loss tensor(1.9158, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 57.422 Acc@5 80.078
^ Epoch: 57
................................Loss tensor(1.7990, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.062 Acc@5 82.812
^ Epoch: 58
................................Loss tensor(1.5733, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.969 Acc@5 84.766
^ Epoch: 59
................................Loss tensor(1.6556, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.844 Acc@5 82.812
^ Epoch: 60
................................Loss tensor(1.6791, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 62.891 Acc@5 84.766
^ Epoch: 61
................................Loss tensor(1.6795, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 65.625 Acc@5 82.812
^ Epoch: 62
................................Loss tensor(1.6260, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.156 Acc@5 84.766
^ Epoch: 63
................................Loss tensor(1.7833, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.453 Acc@5 80.078
^ Epoch: 64
................................Loss tensor(1.5915, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 63.281 Acc@5 84.375
^ Epoch: 65
................................Loss tensor(1.5537, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 63.281 Acc@5 85.547
^ Epoch: 66
................................Loss tensor(1.8186, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 63.281 Acc@5 79.688
^ Epoch: 67
................................Loss tensor(1.7092, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 61.328 Acc@5 84.375
^ Epoch: 68
................................Loss tensor(1.7829, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 62.109 Acc@5 79.688
^ Epoch: 69
................................Loss tensor(1.7229, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 61.719 Acc@5 81.250
^ Epoch: 70
................................Loss tensor(1.8184, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 63.281 Acc@5 81.250
^ Epoch: 71
................................Loss tensor(1.7163, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.844 Acc@5 82.422
^ Epoch: 72
................................Loss tensor(1.4156, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 72.266 Acc@5 86.328
^ Epoch: 73
................................Loss tensor(1.8235, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 59.375 Acc@5 79.688
^ Epoch: 74
................................Loss tensor(1.6704, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 62.109 Acc@5 82.812
^ Epoch: 75
................................Loss tensor(1.6628, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 65.625 Acc@5 82.422
^ Epoch: 76
................................Loss tensor(2.0830, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 52.734 Acc@5 78.906
^ Epoch: 77
................................Loss tensor(1.3728, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.750 Acc@5 85.938
^ Epoch: 78
................................Loss tensor(1.5310, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.016 Acc@5 81.641
^ Epoch: 79
................................Loss tensor(1.6650, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.016 Acc@5 82.812
^ Epoch: 80
................................Loss tensor(1.5176, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.578 Acc@5 83.594
^ Epoch: 81
................................Loss tensor(1.7295, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.453 Acc@5 82.422
^ Epoch: 82
................................Loss tensor(1.6531, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.844 Acc@5 81.641
^ Epoch: 83
................................Loss tensor(1.7386, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.453 Acc@5 81.250
^ Epoch: 84
................................Loss tensor(1.2919, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.531 Acc@5 89.062
^ Epoch: 85
................................Loss tensor(1.6715, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 63.672 Acc@5 83.594
^ Epoch: 86
................................Loss tensor(1.7134, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.844 Acc@5 83.594
^ Epoch: 87
................................Loss tensor(1.3932, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.094 Acc@5 85.938
^ Epoch: 88
................................Loss tensor(1.6647, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.547 Acc@5 84.766
^ Epoch: 89
................................Loss tensor(1.3786, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.094 Acc@5 85.547
^ Epoch: 90
................................Loss tensor(1.8922, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 58.984 Acc@5 79.688
^ Epoch: 91
................................Loss tensor(1.7000, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 65.625 Acc@5 82.422
^ Epoch: 92
................................Loss tensor(1.6285, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 61.719 Acc@5 84.766
^ Epoch: 93
................................Loss tensor(1.4973, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.016 Acc@5 87.109
^ Epoch: 94
................................Loss tensor(1.8492, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.453 Acc@5 80.469
^ Epoch: 95
................................Loss tensor(1.7163, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 62.891 Acc@5 82.812
^ Epoch: 96
................................Loss tensor(1.6591, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.797 Acc@5 81.250
^ Epoch: 97
................................Loss tensor(1.4230, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.141 Acc@5 86.719
^ Epoch: 98
................................Loss tensor(1.5982, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.062 Acc@5 85.938
^ Epoch: 99
................................Loss tensor(1.3960, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.312 Acc@5 84.375
^ Epoch: 100
Epoch 100: Evaluation accuracy on 2048 images, 73.54
................................Loss tensor(1.5578, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.578 Acc@5 83.984
^ Epoch: 101
................................Loss tensor(1.6541, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 63.672 Acc@5 83.984
^ Epoch: 102
................................Loss tensor(1.5631, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.406 Acc@5 83.203
^ Epoch: 103
................................Loss tensor(1.4948, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.797 Acc@5 84.766
^ Epoch: 104
................................Loss tensor(1.5479, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 63.281 Acc@5 84.375
^ Epoch: 105
................................Loss tensor(1.3276, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.312 Acc@5 86.719
^ Epoch: 106
................................Loss tensor(1.2428, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.484 Acc@5 88.672
^ Epoch: 107
................................Loss tensor(1.4942, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.922 Acc@5 83.984
^ Epoch: 108
................................Loss tensor(1.7106, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.547 Acc@5 82.422
^ Epoch: 109
................................Loss tensor(1.5470, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.797 Acc@5 83.984
^ Epoch: 110
................................Loss tensor(1.3573, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.484 Acc@5 87.500
^ Epoch: 111
................................Loss tensor(1.5795, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 65.234 Acc@5 85.156
^ Epoch: 112
................................Loss tensor(1.6091, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 62.500 Acc@5 83.984
^ Epoch: 113
................................Loss tensor(1.6838, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.844 Acc@5 82.031
^ Epoch: 114
................................Loss tensor(1.5490, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.016 Acc@5 83.984
^ Epoch: 115
................................Loss tensor(1.5554, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.453 Acc@5 84.375
^ Epoch: 116
................................Loss tensor(1.3905, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.141 Acc@5 85.156
^ Epoch: 117
................................Loss tensor(1.3837, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.094 Acc@5 86.328
^ Epoch: 118
................................Loss tensor(1.5075, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 65.234 Acc@5 83.203
^ Epoch: 119
................................Loss tensor(1.7448, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.938 Acc@5 81.641
^ Epoch: 120
................................Loss tensor(1.5920, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 65.625 Acc@5 86.328
^ Epoch: 121
................................Loss tensor(1.6689, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 62.891 Acc@5 83.984
^ Epoch: 122
................................Loss tensor(1.5063, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.797 Acc@5 84.375
^ Epoch: 123
................................Loss tensor(1.4001, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.922 Acc@5 86.328
^ Epoch: 124
................................Loss tensor(1.3500, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.922 Acc@5 88.281
^ Epoch: 125
................................Loss tensor(1.3780, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.922 Acc@5 88.672
^ Epoch: 126
................................Loss tensor(1.2107, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.094 Acc@5 88.672
^ Epoch: 127
................................Loss tensor(1.5701, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 61.719 Acc@5 84.766
^ Epoch: 128
................................Loss tensor(1.5812, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.453 Acc@5 82.422
^ Epoch: 129
................................Loss tensor(1.2230, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.094 Acc@5 87.891
^ Epoch: 130
................................Loss tensor(1.3802, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.359 Acc@5 86.328
^ Epoch: 131
................................Loss tensor(1.1879, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.094 Acc@5 88.672
^ Epoch: 132
................................Loss tensor(1.5121, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.797 Acc@5 83.984
^ Epoch: 133
................................Loss tensor(1.6202, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.844 Acc@5 83.594
^ Epoch: 134
................................Loss tensor(1.5210, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.359 Acc@5 86.328
^ Epoch: 135
................................Loss tensor(1.2198, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 72.266 Acc@5 87.500
^ Epoch: 136
................................Loss tensor(1.2879, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.094 Acc@5 87.109
^ Epoch: 137
................................Loss tensor(1.3904, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.094 Acc@5 85.938
^ Epoch: 138
................................Loss tensor(1.4012, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.359 Acc@5 87.109
^ Epoch: 139
................................Loss tensor(1.1848, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.703 Acc@5 88.672
^ Epoch: 140
................................Loss tensor(1.6026, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 63.672 Acc@5 84.766
^ Epoch: 141
................................Loss tensor(1.4604, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.969 Acc@5 87.109
^ Epoch: 142
................................Loss tensor(1.3355, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.016 Acc@5 90.625
^ Epoch: 143
................................Loss tensor(1.4746, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.969 Acc@5 85.938
^ Epoch: 144
................................Loss tensor(1.7988, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.938 Acc@5 83.984
^ Epoch: 145
................................Loss tensor(1.4864, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.750 Acc@5 85.156
^ Epoch: 146
................................Loss tensor(1.4086, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.969 Acc@5 85.938
^ Epoch: 147
................................Loss tensor(1.3260, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.359 Acc@5 89.062
^ Epoch: 148
................................Loss tensor(1.4068, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.578 Acc@5 87.109
^ Epoch: 149
................................Loss tensor(1.3768, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.359 Acc@5 86.719
^ Epoch: 150
Epoch 150: Evaluation accuracy on 2048 images, 76.03
................................Loss tensor(1.5715, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.406 Acc@5 84.766
^ Epoch: 151
................................Loss tensor(1.4310, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.750 Acc@5 85.938
^ Epoch: 152
................................Loss tensor(1.5678, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.750 Acc@5 85.547
^ Epoch: 153
................................Loss tensor(1.6760, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 65.625 Acc@5 84.375
^ Epoch: 154
................................Loss tensor(1.3229, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.969 Acc@5 85.156
^ Epoch: 155
................................Loss tensor(1.4149, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.578 Acc@5 89.062
^ Epoch: 156
................................Loss tensor(1.3266, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 65.625 Acc@5 88.672
^ Epoch: 157
................................Loss tensor(1.3932, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.578 Acc@5 87.500
^ Epoch: 158
................................Loss tensor(1.4213, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.578 Acc@5 86.328
^ Epoch: 159
................................Loss tensor(1.4823, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.797 Acc@5 85.938
^ Epoch: 160
................................Loss tensor(1.4282, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.141 Acc@5 87.109
^ Epoch: 161
................................Loss tensor(1.3808, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.531 Acc@5 87.891
^ Epoch: 162
................................Loss tensor(1.7226, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 63.672 Acc@5 80.469
^ Epoch: 163
................................Loss tensor(1.6555, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 60.156 Acc@5 84.766
^ Epoch: 164
................................Loss tensor(1.5760, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.188 Acc@5 83.984
^ Epoch: 165
................................Loss tensor(1.4708, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.797 Acc@5 89.062
^ Epoch: 166
................................Loss tensor(1.4824, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.359 Acc@5 85.547
^ Epoch: 167
................................Loss tensor(1.4405, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.406 Acc@5 85.156
^ Epoch: 168
................................Loss tensor(1.4300, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.969 Acc@5 85.938
^ Epoch: 169
................................Loss tensor(1.3506, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.922 Acc@5 88.672
^ Epoch: 170
................................Loss tensor(1.2628, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.922 Acc@5 86.719
^ Epoch: 171
................................Loss tensor(1.4226, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.797 Acc@5 84.375
^ Epoch: 172
................................Loss tensor(1.3622, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.750 Acc@5 83.984
^ Epoch: 173
................................Loss tensor(1.2824, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.969 Acc@5 89.844
^ Epoch: 174
................................Loss tensor(1.3369, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.141 Acc@5 89.062
^ Epoch: 175
................................Loss tensor(1.0479, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 76.953 Acc@5 90.234
^ Epoch: 176
................................Loss tensor(1.4030, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.797 Acc@5 86.328
^ Epoch: 177
................................Loss tensor(1.3376, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 65.625 Acc@5 88.281
^ Epoch: 178
................................Loss tensor(1.2719, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.922 Acc@5 87.500
^ Epoch: 179
................................Loss tensor(1.3883, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.312 Acc@5 86.719
^ Epoch: 180
................................Loss tensor(1.4609, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.969 Acc@5 84.375
^ Epoch: 181
................................Loss tensor(1.3293, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.703 Acc@5 85.547
^ Epoch: 182
................................Loss tensor(1.3884, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.531 Acc@5 85.938
^ Epoch: 183
................................Loss tensor(1.4152, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.969 Acc@5 87.891
^ Epoch: 184
................................Loss tensor(1.1994, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.484 Acc@5 90.625
^ Epoch: 185
................................Loss tensor(1.2886, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.922 Acc@5 87.109
^ Epoch: 186
................................Loss tensor(1.3203, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.094 Acc@5 85.938
^ Epoch: 187
................................Loss tensor(1.4532, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.531 Acc@5 85.547
^ Epoch: 188
................................Loss tensor(1.2936, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.922 Acc@5 88.281
^ Epoch: 189
................................Loss tensor(1.4094, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.797 Acc@5 87.109
^ Epoch: 190
................................Loss tensor(1.3561, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.312 Acc@5 85.156
^ Epoch: 191
................................Loss tensor(1.1641, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 73.438 Acc@5 90.234
^ Epoch: 192
................................Loss tensor(1.3716, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.188 Acc@5 87.109
^ Epoch: 193
................................Loss tensor(1.2251, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 72.266 Acc@5 90.234
^ Epoch: 194
................................Loss tensor(1.5889, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 65.625 Acc@5 83.203
^ Epoch: 195
................................Loss tensor(1.2739, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.484 Acc@5 87.500
^ Epoch: 196
................................Loss tensor(1.4433, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.016 Acc@5 85.938
^ Epoch: 197
................................Loss tensor(1.4390, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.141 Acc@5 85.938
^ Epoch: 198
................................Loss tensor(1.3157, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.188 Acc@5 88.672
^ Epoch: 199
................................Loss tensor(1.3140, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.312 Acc@5 86.719
^ Epoch: 200
Epoch 200: Evaluation accuracy on 2048 images, 75.78
................................Loss tensor(1.3315, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.094 Acc@5 85.938
^ Epoch: 201
................................Loss tensor(1.3413, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.484 Acc@5 87.109
^ Epoch: 202
................................Loss tensor(1.4153, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.484 Acc@5 85.938
^ Epoch: 203
................................Loss tensor(1.3798, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.703 Acc@5 87.500
^ Epoch: 204
................................Loss tensor(1.2762, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.922 Acc@5 88.672
^ Epoch: 205
................................Loss tensor(1.2335, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.484 Acc@5 87.500
^ Epoch: 206
................................Loss tensor(1.3859, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.703 Acc@5 85.938
^ Epoch: 207
................................Loss tensor(1.3472, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.359 Acc@5 88.281
^ Epoch: 208
................................Loss tensor(1.4533, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.531 Acc@5 85.938
^ Epoch: 209
................................Loss tensor(1.3123, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.969 Acc@5 87.891
^ Epoch: 210
................................Loss tensor(1.2550, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.875 Acc@5 89.453
^ Epoch: 211
................................Loss tensor(1.5244, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.016 Acc@5 84.375
^ Epoch: 212
................................Loss tensor(1.3878, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.922 Acc@5 85.156
^ Epoch: 213
................................Loss tensor(1.4011, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.531 Acc@5 86.719
^ Epoch: 214
................................Loss tensor(1.3344, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.578 Acc@5 87.109
^ Epoch: 215
................................Loss tensor(1.3703, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.062 Acc@5 87.109
^ Epoch: 216
................................Loss tensor(1.3214, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.750 Acc@5 89.062
^ Epoch: 217
................................Loss tensor(1.3023, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 74.219 Acc@5 85.547
^ Epoch: 218
................................Loss tensor(1.3984, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.578 Acc@5 85.547
^ Epoch: 219
................................Loss tensor(1.3901, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.406 Acc@5 87.891
^ Epoch: 220
................................Loss tensor(1.4091, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 64.844 Acc@5 86.328
^ Epoch: 221
................................Loss tensor(1.4117, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 65.234 Acc@5 85.156
^ Epoch: 222
................................Loss tensor(1.2579, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.094 Acc@5 89.062
^ Epoch: 223
................................Loss tensor(1.4577, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.797 Acc@5 86.719
^ Epoch: 224
................................Loss tensor(1.7113, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 61.719 Acc@5 84.766
^ Epoch: 225
................................Loss tensor(1.1405, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 75.391 Acc@5 88.672
^ Epoch: 226
................................Loss tensor(1.1609, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.875 Acc@5 86.328
^ Epoch: 227
................................Loss tensor(1.1022, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 73.047 Acc@5 89.453
^ Epoch: 228
................................Loss tensor(1.5021, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.578 Acc@5 82.812
^ Epoch: 229
................................Loss tensor(1.1502, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 76.562 Acc@5 89.844
^ Epoch: 230
................................Loss tensor(1.6015, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 63.672 Acc@5 85.547
^ Epoch: 231
................................Loss tensor(1.2069, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.703 Acc@5 89.062
^ Epoch: 232
................................Loss tensor(1.3073, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.312 Acc@5 87.500
^ Epoch: 233
................................Loss tensor(1.3223, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.406 Acc@5 87.500
^ Epoch: 234
................................Loss tensor(1.2075, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.312 Acc@5 90.234
^ Epoch: 235
................................Loss tensor(1.1568, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.094 Acc@5 88.281
^ Epoch: 236
................................Loss tensor(1.3497, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.922 Acc@5 86.328
^ Epoch: 237
................................Loss tensor(1.3879, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.016 Acc@5 88.672
^ Epoch: 238
................................Loss tensor(1.3531, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.312 Acc@5 86.328
^ Epoch: 239
................................Loss tensor(1.3824, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.750 Acc@5 87.109
^ Epoch: 240
................................Loss tensor(1.2173, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 75.391 Acc@5 87.891
^ Epoch: 241
................................Loss tensor(1.4641, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.969 Acc@5 85.156
^ Epoch: 242
................................Loss tensor(1.3383, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.703 Acc@5 87.891
^ Epoch: 243
................................Loss tensor(1.3670, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.875 Acc@5 87.109
^ Epoch: 244
................................Loss tensor(1.2927, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.484 Acc@5 89.062
^ Epoch: 245
................................Loss tensor(1.2431, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 73.438 Acc@5 87.500
^ Epoch: 246
................................Loss tensor(1.1969, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 72.656 Acc@5 88.281
^ Epoch: 247
................................Loss tensor(1.6445, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 62.891 Acc@5 84.375
^ Epoch: 248
................................Loss tensor(1.2642, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 68.359 Acc@5 89.453
^ Epoch: 249
................................Loss tensor(1.3059, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 73.438 Acc@5 85.547
^ Epoch: 250
Epoch 250: Evaluation accuracy on 2048 images, 75.00
................................Loss tensor(1.0560, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 73.047 Acc@5 92.188
^ Epoch: 251
................................Loss tensor(1.1369, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.531 Acc@5 87.891
^ Epoch: 252
................................Loss tensor(1.3569, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.312 Acc@5 87.891
^ Epoch: 253
................................Loss tensor(1.2233, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 73.438 Acc@5 89.062
^ Epoch: 254
................................Loss tensor(1.2344, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 72.656 Acc@5 90.625
^ Epoch: 255
................................Loss tensor(1.5015, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 65.234 Acc@5 83.203
^ Epoch: 256
................................Loss tensor(1.0261, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 78.516 Acc@5 90.625
^ Epoch: 257
................................Loss tensor(1.2063, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.312 Acc@5 90.234
^ Epoch: 258
................................Loss tensor(1.1763, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.703 Acc@5 89.453
^ Epoch: 259
................................Loss tensor(1.3709, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 67.969 Acc@5 87.891
^ Epoch: 260
................................Loss tensor(1.1821, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 72.656 Acc@5 89.062
^ Epoch: 261
................................Loss tensor(1.3488, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 66.406 Acc@5 88.281
^ Epoch: 262
................................Loss tensor(1.1726, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 72.266 Acc@5 89.453
^ Epoch: 263
................................Loss tensor(1.3029, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.312 Acc@5 88.672
^ Epoch: 264
................................Loss tensor(1.2436, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.922 Acc@5 89.844
^ Epoch: 265
................................Loss tensor(1.3371, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 69.531 Acc@5 86.328
^ Epoch: 266
................................Loss tensor(1.3199, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.312 Acc@5 87.891
^ Epoch: 267
................................Loss tensor(1.2755, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.094 Acc@5 88.672
^ Epoch: 268
................................Loss tensor(1.2874, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 71.094 Acc@5 87.500
^ Epoch: 269
................................Loss tensor(1.1613, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 72.656 Acc@5 90.234
^ Epoch: 270
................................Loss tensor(1.0151, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 77.344 Acc@5 89.844
^ Epoch: 271
................................Loss tensor(1.2258, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.703 Acc@5 88.672
^ Epoch: 272
................................Loss tensor(1.1449, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 75.781 Acc@5 89.453
^ Epoch: 273
................................Loss tensor(1.0311, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 77.344 Acc@5 90.625
^ Epoch: 274
................................Loss tensor(1.2737, device='cuda:0', grad_fn=<DivBackward0>)
Training: * Acc@1 70.703 Acc@5 89.453
^ Epoch: 275

6. ONNX export and post-processing

We use torch.onnx.dynamo_export() to export the graph from PyTorch to ONNX and use Quadric's QAT post-processor in order to clean up the graph and format it nicely. We start here with the conversion from PyTorch to ONNX.

<div class="alert alert-block alert-warning"> <b>Note:</b> Because torch.onnx.dynamo_export is a WIP, there will be many warnings generated for now. There are no adverse effects. These are just remnants from other functionalities, so they can be ignored. But they cannot be suppressed as they are thrown by the underlying C++. </div>

if not pretrained_mobilenet:
    # Workaround for a bug in convert_pt2e where the function would fail
    # to execute due to improperly handled tensors when GPUs are used
    torch.set_default_device("cpu")
    prepared_model = prepared_model.to("cpu")
    prepared_model.recompile()
    quantized_model = convert_pt2e(prepared_model)
    torch.set_default_device("cuda")
    quantized_model = quantized_model.to("cuda")
    quantized_model.recompile()

    torch.ao.quantization.move_exported_model_to_eval(quantized_model)
    top1, top5 = evaluate(
        quantized_model, data_loader_test, device="cuda", neval_batches=num_eval_batches
    )
    print(
        "Final evaluation accuracy on %d images, (Top-1): %2.2f"
        % (num_eval_batches * test_batch_size, top1.avg)
    )
    print(
        "Final evaluation accuracy on %d images, (Top-5): %2.2f"
        % (num_eval_batches * test_batch_size, top5.avg)
    )

    onnx_program_file = tempfile.NamedTemporaryFile()
    onnx_program = torch.onnx.dynamo_export(quantized_model, example_input)
    onnx_program.save(onnx_program_file.name)
else:
    onnx_program_file = tempfile.NamedTemporaryFile()
    torch.onnx.export(
        float_model,
        example_input,
        onnx_program_file.name,
        export_params=True,
        opset_version=16,
        do_constant_folding=True,
        input_names=["input"],
    )
Final evaluation accuracy on 2048 images, (Top-1): 75.93
Final evaluation accuracy on 2048 images, (Top-5): 93.26


/usr/local/lib/python3.10/dist-packages/torch/onnx/_internal/_exporter_legacy.py:116: UserWarning: torch.onnx.dynamo_export only implements opset version 18 for now. If you need to use a different opset version, please register them with register_custom_op.
  warnings.warn(
/usr/local/lib/python3.10/dist-packages/torch/onnx/_internal/fx/passes/readability.py:52: UserWarning: Attempted to insert a get_attr Node with no underlying reference in the owning GraphModule! Call GraphModule.add_submodule to add the necessary submodule, GraphModule.add_parameter to add the necessary Parameter, or nn.Module.register_buffer to add the necessary buffer
  new_node = self.module.graph.get_attr(normalized_name)
/usr/local/lib/python3.10/dist-packages/torch/fx/graph.py:1586: UserWarning: Node features_0_1_running_var target features_0_1_running_var features_0_1_running_var of  does not reference an nn.Module, nn.Parameter, or buffer, which is what 'get_attr' Nodes typically target
  warnings.warn(f'Node {node} target {node.target} {atom} of {seen_qualname} does '
/usr/local/lib/python3.10/dist-packages/torch/fx/graph.py:1586: UserWarning: Node _frozen_param0 target _frozen_param0 _frozen_param0 of  does not reference an nn.Module, nn.Parameter, or buffer, which is what 'get_attr' Nodes typically target
  warnings.warn(f'Node {node} target {node.target} {atom} of {seen_qualname} does '
/usr/local/lib/python3.10/dist-packages/torch/fx/graph.py:1586: UserWarning: Node features_0_1_running_mean target features_0_1_running_mean features_0_1_running_mean of  does not reference an nn.Module, nn.Parameter, or buffer, which is what 'get_attr' Nodes typically target
  warnings.warn(f'Node {node} target {node.target} {atom} of {seen_qualname} does '
/usr/local/lib/python3.10/dist-packages/torch/fx/graph.py:1586: UserWarning: Node features_0_1_running_var_1 target features_0_1_running_var features_0_1_running_var of  does not reference an nn.Module, nn.Parameter, or buffer, which is what 'get_attr' Nodes typically target
  warnings.warn(f'Node {node} target {node.target} {atom} of {seen_qualname} does '
/usr/local/lib/python3.10/dist-packages/torch/fx/graph.py:1586: UserWarning: Node features_1_0_block_0_1_running_var target features_1_0_block_0_1_running_var features_1_0_block_0_1_running_var of  does not reference an nn.Module, nn.Parameter, or buffer, which is what 'get_attr' Nodes typically target
  warnings.warn(f'Node {node} target {node.target} {atom} of {seen_qualname} does '
/usr/local/lib/python3.10/dist-packages/torch/fx/graph.py:1593: UserWarning: Additional 758 warnings suppressed about get_attr references
  warnings.warn(
/usr/local/lib/python3.10/dist-packages/torch/onnx/_internal/fx/onnxfunction_dispatcher.py:503: FutureWarning: 'onnxscript.values.TracedOnnxFunction.param_schemas' is deprecated in version 0.1 and will be removed in the future. Please use '.op_signature' instead.
  self.param_schema = self.onnxfunction.param_schemas()

The next step uses Quadric's PyTorch graph post-processor. Due to operator and network dependent outputs from PyTorch, the post-processor may not yet support the graph if you have changed the model being quantized.

Supported ops/patterns:

  • Conv (+ BatchNormalization)
  • Gemm
  • GlobalAveragePool
  • Add (i.e. residual adds)
  • aten_bernoulli_p PyTorch dropout layers

Unsupported ops/patterns:

  • ATen operators from PyTorch aside from aten_bernoulli_p
torch.set_default_device("cpu")
output_path = (
    "./efficientnet-processed.onnx" if not pretrained_mobilenet else "./mobilenet-processed.onnx"
)
qat_processor = processor.QATProcessor()
qat_processor.process(
    onnx_program_file.name,
    q_add_flag=True,
    strip_onnx=False,
    output_path=output_path,
)
2026-07-18 13:56 - DEBUG - epu - qat_processor - Preparing calibration data
/usr/local/lib/python3.10/dist-packages/datasets/table.py:1421: FutureWarning: promote has been superseded by promote_options='default'.
  table = cls._concat_blocks(blocks, axis=0)
2026-07-18 13:56 - DEBUG - epu - qat_processor - Fixing PyTorch exported Unsqueeze axes arg
2026-07-18 13:56 - DEBUG - epu - qat_processor - Removing aten.bernoulli ops
2026-07-18 13:56 - DEBUG - epu - qat_processor - Removing NOP aten.pad ops
2026-07-18 13:56 - DEBUG - epu - qat_processor - Removing aten.as_strided ops
2026-07-18 13:56 - DEBUG - epu - qat_processor - Swapping aten.gelu ops for ONNX equivalent
2026-07-18 13:56 - DEBUG - epu - qat_processor - Swapping aten.roll ops for ONNX equivalent
2026-07-18 13:57 - DEBUG - epu - qat_processor - Swapping ReduceMean for GlobalAveragePool
2026-07-18 13:57 - DEBUG - epu - qat_processor - Removing Dropout ops with rate=0
2026-07-18 13:57 - DEBUG - epu - qat_processor - Removing redundant QDQs
2026-07-18 13:57 - DEBUG - epu - qat_processor - Folding Div into BatchNormalization
2026-07-18 13:57 - DEBUG - epu - qat_processor - Folding BatchNormalization into Conv nodes with biases
2026-07-18 13:57 - DEBUG - epu - qat_processor - Folding BatchNormalization into Conv nodes without biases
2026-07-18 13:57 - DEBUG - epu - qat_processor - Quantizing Swin MHA blocks
2026-07-18 13:57 - DEBUG - epu - qat_processor - Quantizing FC layers
2026-07-18 13:57 - DEBUG - epu - qat_processor - Quantizing QKV MatMul ops
2026-07-18 13:57 - DEBUG - epu - qat_processor - Quantizing residual adds
WARNING:root:Please use QuantFormat.QDQ for activation type QInt8 and weight type QInt8. Or it will lead to bad performance on x64.
WARNING:root:Please check if the model is already quantized. Note you don't need to quantize a QAT model. OnnxRuntime support to run QAT model directly.
2026-07-18 13:58 - DEBUG - epu - qat_processor - Changing to Opset 17
2026-07-18 13:58 - DEBUG - epu - qat_processor - Performing ORT optimizations
2026-07-18 13:58 - DEBUG - epu - qat_processor - Unfolding QLinearSoftmax to avoid inaccurate LUT
2026-07-18 13:58 - DEBUG - epu - qat_processor - Fusing QLinearConv
2026-07-18 13:58 - DEBUG - epu - qat_processor - Fusing QGemm
2026-07-18 13:58 - DEBUG - epu - qat_processor - Folding DQ into MatMul ops where ORT does not
2026-07-18 13:58 - DEBUG - epu - qat_processor - Moving DQ outside of Roll patterns
2026-07-18 13:58 - DEBUG - epu - qat_processor - Moving from uint8 to int8

7. Prepare the validation dataset

Now that the QAT model has been saved to disk, we will compare it to the original pretrained FP32 EfficientNet-B7 model from PyTorch. Here we use a set of 3500 images in our validation of the model.

torch.set_default_device("cpu")
MAX_NUM_SAMPLES_FOR_ACCURACY_CALCULATION = 3500
if not pretrained_mobilenet:
    graph_input_name = onnx_program.model_proto.graph.input[0].name
else:
    graph_input_name = "input"
assert isinstance(graph_input_name, str)

test_transforms = Compose(
    [
        Resize(224),
        CenterCrop(224),
        ToTensor(),
        normalize,
    ]
)

dataset = ImageNet_Mini_Quadric.Dataset(transform=test_transforms)
subset_of_dataset = Subset(dataset, range(MAX_NUM_SAMPLES_FOR_ACCURACY_CALCULATION))
calibration_dataloader = CalibrationDataLoader(
    DataLoader(subset_of_dataset, batch_size=1, shuffle=True), [graph_input_name]
)

8. Validate and compare against FP32 model

Now we want to make sure that quantization was done succesfully by checking the performance of the original model against the quantized model. If you skipped the training step, feel free to use pretrained_qat_path instead of int8_path when initializing quantized_onnx_model.

Depending on the task, you may choose different metrics like spearman correlation, f1-score, AUC, ROC, L2 norm, IoU etc.

if not pretrained_mobilenet:
    pytorch_model = efficientnet_b7(
        weights=EfficientNet_B7_Weights.DEFAULT,
    )
else:
    pytorch_model = mobilenet_v2(
        weights=MobileNet_V2_Weights.IMAGENET1K_V1,
    )
int8_path = output_path
fp32_path = f"./{pytorch_model.__class__.__name__}_float32.onnx"

## NOTE: uncomment the 2 lines below to use our included pretrained EfficientNet file
## int8_path = "../models/efficientnet/efficientnet-quadric-qat.onnx"
## graph_input_name = "l_x_"

onnx_model_path = Path(fp32_path)
quantized_onnx_model = QuantizedONNXModel(int8_path, None)
input_tensor_shape = (1, 3, 224, 224)
example_input = torch.randn(*input_tensor_shape, requires_grad=True)
torch.onnx.export(
    pytorch_model,
    example_input,
    str(onnx_model_path),
    export_params=True,
    do_constant_folding=True,
    opset_version=DEFAULT_ONNX_OPSET,
    input_names=[graph_input_name],
    output_names=["output"],
)

## Evaluation
quantization_experiment = QuantizationExperiment(
    floating_point_onnx_model_path=onnx_model_path,
    quantized_onnx_model=quantized_onnx_model,
    performance_tracker=ClassifierPerformanceTracker(),
)
run_quantization_experiment(
    quantization_experiment,
    calibration_dataloader,
    export_path=None,
    max_num_samples=MAX_NUM_SAMPLES_FOR_ACCURACY_CALCULATION,
)
3500: FP3274.26% <> INT873.51%: 100%|██████████| 3500/3500 [15:37<00:00,  3.74it/s]


Used 3500 image samples for model accuracy comparison.
Original FP32 model accuracy (Top-1): 74.26%
Quantized INT8 model accuracy (Top-1): 73.51%
Change in Top-1 model accuracy due to quantization: 0.74%

9. CGC lowering (pretrained MobileNet only)

For now, this step is only supported for the QAT pretrained MobileNet, not the QAT EfficientNet.

if pretrained_mobilenet:
    input_image_path = "../common/calibration/raw/millie_imagenet.jpeg"
    input_image = Image.open(input_image_path)
    transformed_image = test_transforms(input_image)

    hw_config = HWConfig(ocm_size="4MB")
    cgc_job = ChimeraJob(
        model_p=int8_path,
        hw_config=hw_config,
        validate_iss=True,
    )
    cgc_job.compile()
    cgc_job.validate_ort_iss(inputs={"input": np.expand_dims(transformed_image.numpy(), 0)})

Table of Contents
Introduction to the Chimera SDK
Chimera SDK Quick Start Guide
Chimera SDK Command Line Interface (CLI)
Tutorial: Using SDK as a Library
Tutorials & Model Demos
Quantization Tutorials
Multicore Demo
Chimera LLVM C++ Compiler
Chimera SDK Licensing Policy Documentation
Glossary

Sign in to your account

Don't have an account? Create an Account
By signing in, you are agreeing to our Terms of Use and Privacy Policy.

Develop.

Simulate.

Profile.

Collaborate.