Reference for ultralytics/utils/export/litert.py#
This page is sourced from https://github.com/ultralytics/ultralytics/blob/main/ultralytics/utils/export/litert.py. Have an improvement or example to add? Open a Pull Request — thank you! 🙏
Function ultralytics.utils.export.litert._litert_grouped_topk#
def _litert_grouped_topk(x: torch.Tensor, k: int, groups: int) -> tuple[torch.Tensor, torch.Tensor]Select the top k of x along dim 1 with int32 indices, which GPU delegates accept and int64 they do not.
Args
| Name | Type | Description | Default |
|---|---|---|---|
x | torch.Tensor | required | |
k | int | required | |
groups | int | required |
ultralytics/utils/export/litert.py
def _litert_grouped_topk(x: torch.Tensor, k: int, groups: int) -> tuple[torch.Tensor, torch.Tensor]:
"""Select the top k of x along dim 1 with int32 indices, which GPU delegates accept and int64 they do not."""
values, index = Detect._grouped_topk(x, k, groups)
return values, index.int()Function ultralytics.utils.export.litert._litert_gather#
def _litert_gather(self, x: torch.Tensor, index: torch.Tensor) -> torch.TensorSelect index (batch, k) rows of x along dim 1 without gather_nd, which GPU delegates do not implement.
Args
| Name | Type | Description | Default |
|---|---|---|---|
self | required | ||
x | torch.Tensor | required | |
index | torch.Tensor | required |
ultralytics/utils/export/litert.py
def _litert_gather(self, x: torch.Tensor, index: torch.Tensor) -> torch.Tensor:
"""Select index (batch, k) rows of x along dim 1 without gather_nd, which GPU delegates do not implement."""
b, n = x.shape[:2]
offset = torch.arange(b, device=x.device, dtype=index.dtype)[..., None] * n
return x.flatten(0, 1).index_select(0, (index + offset).flatten()).view(b, index.shape[1], *x.shape[2:])Function ultralytics.utils.export.litert.torch2litert#
def torch2litert(
model: torch.nn.Module,
im: torch.Tensor,
file: Path,
quantize: int | str | None,
calibration_dataset: torch.utils.data.DataLoader | None,
metadata: dict | None,
prefix: str,
) -> PathExport a PyTorch model to LiteRT format using litert_torch, with optional INT8 quantization.
Three INT8 schemes are supported via quantize: 8 applies static INT8 (int8 weights + int8 activations) and 'w8a16' applies static INT8 weights with int16 activations, both requiring a calibration_dataset; 'w8a32' applies dynamic/weight-only INT8 (int8 weights + FP32 activations) and needs no calibration. None/32 exports FP32. FP16 is not exported as a separate model: LiteRT runs an FP32 model in FP16 at runtime via the GPU delegate (FP16 by default) or the XNNPACK FORCE_FP16 flag on ARM.
Args
| Name | Type | Description | Default |
|---|---|---|---|
model | torch.nn.Module | The PyTorch model to export. | required |
im | torch.Tensor | Example input tensor for tracing. | required |
file | Path | str | Source model file path used to derive output directory. | required |
quantize | int | str | None | Quantization scheme: 8 (static INT8), 'w8a16' (static int8 weights + int16 activations), 'w8a32' (dynamic INT8), or None/32 (FP32). | required |
calibration_dataset | DataLoader | None | Calibration dataloader for static quantization, as returned by get_int8_calibration_dataloader. Required when quantize is 8 or 'w8a16'. | required |
metadata | dict | None | Optional metadata embedded in the .tflite as a metadata.json entry. | required |
prefix | str | Prefix for log messages. | required |
Returns
| Type | Description |
|---|---|
Path | Path to the exported .tflite file with metadata embedded as a metadata.json entry. |
ultralytics/utils/export/litert.py
def torch2litert(
model: torch.nn.Module,
im: torch.Tensor,
file: Path,
quantize: int | str | None,
calibration_dataset: torch.utils.data.DataLoader | None,
metadata: dict | None,
prefix: str,
) -> Path:
"""Export a PyTorch model to LiteRT format using litert_torch, with optional INT8 quantization.
Three INT8 schemes are supported via ``quantize``: ``8`` applies static INT8 (int8 weights + int8 activations) and
``'w8a16'`` applies static INT8 weights with int16 activations, both requiring a ``calibration_dataset``;
``'w8a32'`` applies dynamic/weight-only INT8 (int8 weights + FP32 activations) and needs no calibration.
``None``/``32`` exports FP32. FP16 is not exported as a separate model: LiteRT runs an FP32 model in FP16 at runtime
via the GPU delegate (FP16 by default) or the XNNPACK ``FORCE_FP16`` flag on ARM.
Args:
model (torch.nn.Module): The PyTorch model to export.
im (torch.Tensor): Example input tensor for tracing.
file (Path | str): Source model file path used to derive output directory.
quantize (int | str | None): Quantization scheme: ``8`` (static INT8), ``'w8a16'`` (static int8 weights + int16
activations), ``'w8a32'`` (dynamic INT8), or ``None``/``32`` (FP32).
calibration_dataset (DataLoader | None): Calibration dataloader for static quantization, as returned by
``get_int8_calibration_dataloader``. Required when ``quantize`` is ``8`` or ``'w8a16'``.
metadata (dict | None): Optional metadata embedded in the ``.tflite`` as a ``metadata.json`` entry.
prefix (str): Prefix for log messages.
Returns:
(Path): Path to the exported ``.tflite`` file with metadata embedded as a ``metadata.json`` entry.
"""
from ultralytics.utils.checks import check_requirements
check_requirements(("litert-torch>=0.9.0", "ai-edge-litert>=2.1.4"))
import litert_torch
static_int8 = quantize == 8
static_int16 = quantize == "w8a16"
dynamic_int8 = quantize == "w8a32"
LOGGER.info(f"\n{prefix} starting export with litert_torch {litert_torch.__version__}...")
file = Path(file)
quant_tag = "_int8" if static_int8 else "_w8a16" if static_int16 else "_w8a32" if dynamic_int8 else ""
# Normalize coordinate channels by input size so INT8 quantization preserves scores (denormalized in LiteRTBackend).
# End-to-end models output post-NMS pixel coordinates in FP32 (no scale collapse), so they are left as-is.
meta = metadata or {}
task = meta.get("task")
if task in {"detect", "segment", "pose", "obb"} and not meta.get("end2end", False):
model = _NormalizeCoords(
model, int(im.shape[2]), int(im.shape[3]), task, len(meta.get("names", {})), meta.get("kpt_shape")
)
for m in model.modules(): # int32 indices and a gather_nd-free gather keep the head on the GPU delegate
if isinstance(m, Detect):
m._grouped_topk = _litert_grouped_topk
m._gather = types.MethodType(_litert_gather, m)
# Lower index_select to tfl.gather: the default lowering emits GATHER_ND, which GPU delegates do not implement
litert_torch.fx_infra.decomp.add_pre_lower_decomp(
torch.ops.aten.index_select.default, lambda x, dim, index: torch.ops.tfl.gather(x, index.int(), dim)
)
edge_model = litert_torch.convert(model, (im,))
tflite_file = file.with_name(f"{file.stem}{quant_tag}.tflite")
edge_model.export(tflite_file)
if static_int8 or static_int16 or dynamic_int8:
check_requirements("ai-edge-quantizer>=0.6.0")
from ai_edge_quantizer import qtyping, quantizer, recipe
qt = quantizer.Quantizer(str(tflite_file))
if static_int8 or static_int16: # static schemes calibrate over representative images
act = "int8" if static_int8 else "int16"
LOGGER.info(f"{prefix} applying static quantization (int8 weights + {act} activations)...")
calib_samples = []
for batch in calibration_dataset:
imgs = batch["img"].cpu().float() / 255.0
# litert-torch traces a fixed batch; tile under-sized batches up to im's batch dim (repeats are
# statistics-identical for calibration)
if imgs.shape[0] < im.shape[0]:
imgs = imgs.repeat(-(-im.shape[0] // imgs.shape[0]), 1, 1, 1)[: im.shape[0]]
calib_samples.append({"args_0": imgs.numpy()})
qt.load_quantization_recipe(recipe.static_wi8_ai8() if static_int8 else recipe.static_wi8_ai16())
# Keep FP32 graph input/output (weights/activations stay int8/int16 internally): matches the historical
# onnx2tf "fp32 in/out" contract that downstream consumers (LiteRT GPU delegate, on-device runtimes) expect,
# and avoids forcing every consumer to (de)quantize at the boundary. Must run after load_quantization_recipe.
for op in (qtyping.TFLOperationName.INPUT, qtyping.TFLOperationName.OUTPUT):
qt.update_quantization_recipe(
regex=".*", operation_name=op, algorithm_key=recipe.AlgorithmName.NO_QUANTIZE
)
result = qt.calibrate({"serving_default": calib_samples})
qt.quantize(calibration_result=result).export_model(str(tflite_file), overwrite=True)
else: # dynamic / weight-only INT8: int8 weights, FP32 activations, no calibration needed
LOGGER.info(f"{prefix} applying dynamic INT8 quantization (int8 weights + FP32 activations)...")
qt.load_quantization_recipe(recipe.dynamic_wi8_afp32())
qt.quantize().export_model(str(tflite_file), overwrite=True)
# Embed metadata as a JSON entry appended to the .tflite (zip-tolerant flatbuffer), so the model is a single
# self-contained file that LiteRTBackend reads back at load time.
with zipfile.ZipFile(tflite_file, "a", zipfile.ZIP_DEFLATED) as zf:
zf.writestr("metadata.json", json.dumps(metadata or {}))
return tflite_file