Note
Go to the end to download the full example code.
Mixed precision search#
Post-training quantization is not all-or-nothing: leaving a few sensitive
layers at full precision often recovers most of the accuracy loss for a small
cost in latency or size. embedl_deploy’s search subsystem automates that
trade-off – it drives an Optuna study over which
activations and weights to skip, scoring each candidate with a user-supplied
evaluation function.
Two strategies ship today:
GranularQuantConfigSearchSpace– one independent boolean per quantizer axis, the finest granularity.PartitionedQuantConfigSearchSpace– a handful of continuous parameters carve the model into alternating skip/keep sections, so the search stays constant-size regardless of model size.
This tutorial takes an ImageNet-pretrained torchvision ResNet-18, fuses it with the TensorRT patterns, and searches for which quantizers to skip – trading ImageNette validation accuracy against how much of the model stays quantized.
Install dependencies:
pip install "embedl-deploy[tensorrt,search]" torchvision
Setup: a pretrained model and ImageNette#
We use an ImageNet-pretrained ResNet-18 as-is – no fine-tuning needed. For data we use ImageNette, a 10-class subset of ImageNet, as in the other tutorials. Its labels are remapped to the corresponding ImageNet class indices so the pretrained classifier head can be scored directly.
import itertools
import tarfile
import urllib.request
from pathlib import Path
import torch
import torchvision
from torchvision import transforms
from torchvision.models import resnet18
from embedl_deploy.backend import set_backend
set_backend("tensorrt")
torch.manual_seed(0)
IMAGENETTE_URL = (
"https://s3.amazonaws.com/fast-ai-imageclas/imagenette2-320.tgz"
)
DATA_DIR = Path("artifacts/data")
IMAGENETTE_DIR = DATA_DIR / "imagenette2-320"
IMAGENETTE_TO_IMAGENET = [0, 217, 482, 491, 497, 566, 569, 571, 574, 701]
BATCH_SIZE = 32
def download_imagenette() -> None:
"""Download and extract ImageNette if not already present."""
if IMAGENETTE_DIR.exists():
print(f"ImageNette already present at {IMAGENETTE_DIR}")
return
DATA_DIR.mkdir(parents=True, exist_ok=True)
tgz_path = DATA_DIR / "imagenette2-320.tgz"
print(f"Downloading ImageNette to {tgz_path} ...")
urllib.request.urlretrieve(IMAGENETTE_URL, str(tgz_path))
print("Extracting ...")
with tarfile.open(tgz_path) as tar:
tar.extractall(DATA_DIR)
tgz_path.unlink()
print("Done.")
download_imagenette()
transform_images = transforms.Compose(
[
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),
]
)
def remap(label: int) -> int:
"""Map an ImageNette class index to its ImageNet class index."""
return IMAGENETTE_TO_IMAGENET[label]
train_set = torchvision.datasets.ImageFolder(
str(IMAGENETTE_DIR / "train"),
transform=transform_images,
target_transform=remap,
)
val_set = torchvision.datasets.ImageFolder(
str(IMAGENETTE_DIR / "val"),
transform=transform_images,
target_transform=remap,
)
train_loader = torch.utils.data.DataLoader(
train_set, batch_size=BATCH_SIZE, shuffle=True
)
model = resnet18(weights="IMAGENET1K_V1").eval()
Held-out evaluation and calibration data#
A fixed slice of the validation set scores every candidate on the same
images, and a few training batches calibrate the quant stubs. Small,
fixed subsets keep each search trial cheap. ImageFolder yields its
samples in class order, so the slice is drawn through a seeded shuffle –
taken straight off the loader it would be one class of the ten.
eval_loader = torch.utils.data.DataLoader(
val_set,
batch_size=BATCH_SIZE,
shuffle=True,
generator=torch.Generator().manual_seed(0),
)
eval_batches = list(itertools.islice(iter(eval_loader), 8))
calib_batches = [
images for images, _ in itertools.islice(iter(train_loader), 4)
]
def accuracy(model: torch.nn.Module) -> float:
"""Top-1 accuracy of `model` on the fixed evaluation batches."""
correct = total = 0
with torch.no_grad():
for images, labels in eval_batches:
correct += int((model(images).argmax(dim=1) == labels).sum())
total += labels.numel()
return correct / total
fp32_accuracy = accuracy(model)
print(f"FP32 accuracy: {fp32_accuracy:.1%}")
Fusing for quantization#
The search operates on a fused, not-yet-quantized model.
from embedl_deploy import transform
from embedl_deploy.tensorrt import TENSORRT_PATTERNS
args = (torch.randn(1, 3, 224, 224),)
fused = transform(model, args, patterns=TENSORRT_PATTERNS).model
print(f"Fused accuracy: {accuracy(fused):.1%} (fusion is lossless)")
Scoring a configuration#
eval_fn receives a freshly quantized model and returns any number of
named metrics. Here, accuracy is the real test-subset accuracy, and
active_quantizers counts how many quant stubs are still enabled –
fewer skips means more of the model runs in INT8 but risks more accuracy
loss, so the two trade off against each other as the search progresses.
from embedl_deploy.quantize import QuantStub, WeightFakeQuantize
from embedl_deploy.search import Direction
def forward_loop(model: torch.nn.Module) -> None:
"""Feed the calibration batches through `model` to calibrate its stubs."""
with torch.no_grad():
for images in calib_batches:
model(images)
def eval_fn(model: torch.nn.Module) -> dict[str, float]:
"""Score `model` on accuracy and on how many quantizers stay active."""
active = sum(
module.enabled
for module in model.modules()
if isinstance(module, (QuantStub, WeightFakeQuantize))
)
return {"accuracy": accuracy(model), "active_quantizers": active}
directions: dict[str, Direction] = {
"accuracy": "maximize",
"active_quantizers": "maximize",
}
The search strategies#
Both strategies discover the same underlying dimensions – one per enabled weight or activation quantizer – from the fused model alone, before any search-space decision is made. Each dimension is keyed by the module and axis it controls, and those keys double as the Optuna parameter names, so trial parameters stay readable in model terms everywhere downstream.
from embedl_deploy.search import get_dimensions
dimensions = get_dimensions(fused)
print(f"{len(dimensions)} quantizable dimensions, e.g.:")
for key in list(dimensions)[:4]:
print(f" {key}")
Running a granular search#
GranularQuantConfigSearchSpace turns every dimension into its own boolean
parameter.
search()
drives Optuna’s ask/evaluate/tell loop, seeding two informed starting
points ahead of the sampled ones: skip nothing (the fully-INT8 model) and
skip everything (effectively the FP32 baseline).
from embedl_deploy.search import GranularQuantConfigSearchSpace
granular = GranularQuantConfigSearchSpace(
fused, args, forward_loop=forward_loop, eval_fn=eval_fn
)
results, study = granular.search(16, directions=directions)
int8_result = results[0]
if int8_result is None:
raise RuntimeError("The seeded fully-INT8 trial failed.")
int8_accuracy = int8_result["accuracy"]
print(f"FP32 baseline: {fp32_accuracy:.1%}")
print(f"Fully INT8: {int8_accuracy:.1%}")
print("Pareto-optimal trade-offs found:")
for trial in sorted(study.best_trials, key=lambda t: t.values[1]):
active = int(trial.values[1])
print(f" {active:2d} active quantizers -> {trial.values[0]:.1%}")
Inspecting a trial#
Because parameters are named after the modules they control, a trial’s
params read directly in model terms. Each trial additionally records
the concrete manifestation of its point – the qualified names it skips,
per axis – on its skip user attribute, so Optuna tooling
(study.trials_dataframe(), dashboards, plot hover text) shows what
a trial did, not just its raw parameters.
trial = study.trials[2]
skipped = [key for key, keep in trial.params.items() if not keep]
print(f"trial 2 skips {len(skipped)} axes, e.g. {skipped[:3]}")
print(f"manifestation: {trial.user_attrs['skip']}")
print(f"as a config: {granular.skip(trial)!r}")
Constant-size search spaces#
For a model with hundreds of quantizable dimensions, one boolean each is
a lot of parameters. PartitionedQuantConfigSearchSpace instead carves the
same dimensions into a handful of alternating skip/keep sections,
controlled by only 2 * num_enabled_sections continuous parameters –
constant-size regardless of model size. Its raw parameters are abstract
cut positions, which is exactly where the per-trial skip user
attribute keeps the study interpretable.
from embedl_deploy.search import PartitionedQuantConfigSearchSpace
partitioned = PartitionedQuantConfigSearchSpace(
fused, args, forward_loop=forward_loop, eval_fn=eval_fn
)
print(partitioned)
Crash-resilient persistence#
Passing path persists every trial via
~optuna.storages.JournalStorage – a crashed run resumes exactly where
it left off instead of re-running completed trials. The journal lives
under artifacts/ next to the data, so re-running this script (or
raising n_trials later) picks the study up where it stopped rather
than starting over.
STUDY_PATH = Path("artifacts/search/partitioned_study.log")
STUDY_PATH.parent.mkdir(parents=True, exist_ok=True)
first, _ = partitioned.search(
4, directions=directions, path=STUDY_PATH, name="partitioned"
)
second, _ = partitioned.search(
len(first) + 4, directions=directions, path=STUDY_PATH, name="partitioned"
)
if second[: len(first)] != first:
raise RuntimeError("Resuming re-ran trials that `path` already held.")
print(
f"resumed run reused {len(first)} prior trials, "
f"added {len(second) - len(first)} more"
)
Rebuilding the best model#
With two objectives there is no single best trial, only a Pareto front
(study.best_trials). Pick the one you want by your own rule – here the
most-quantized trial that stays within one point of the FP32 accuracy –
and hand it back to the space. config returns the full QuantConfig
that trial was scored with, and apply re-quantizes a fresh copy of the
fused model with it, ready for export. Both accept a FrozenTrial from a
study reloaded via search(0, path=..., name=...) in a later session,
so the search never has to be re-run to get its model back.
within_budget = [
t for t in study.best_trials if t.values[0] >= fp32_accuracy - 0.01
]
best = max(within_budget or study.best_trials, key=lambda t: t.values[1])
print(f"best trial {best.number}: {best.user_attrs['result']}")
print(f"its config: {granular.config(best)}")
best_model = granular.apply(best)
print(f"rebuilt accuracy: {accuracy(best_model):.1%}")
Visualizing the trade-off#
If Plotly is installed then Optuna’s own visualization utilities work
unchanged, since search just drives a normal optuna.Study.
Hovering a point shows the trial’s skip user attribute – the modules
that configuration leaves at full precision.
import optuna.visualization
fig = optuna.visualization.plot_pareto_front(
study,
target_names=["active_quantizers", "accuracy"],
targets=lambda t: (t.values[1], t.values[0]),
)
fig.write_html("pareto_front.html")
print("Pareto front written to pareto_front.html")