Summary
With get_symmetric_a16w8_quantization_config, aten.rsqrt gets lowered by InsertTableOpsPass to an int16 TOSA TABLE. The table holds 513 points spread evenly over the whole int16 code range of the input, and TOSA interpolates linearly between them. rsqrt is steep near zero, so the first segments are badly wrong: for inputs below about 1/250 of the calibrated input max, the output comes out up to ~2× too large in the minimal repro below, and up to 5.7× in a trained model. PT2E fake-quant computes an exact rsqrt between Q/DQ nodes, so it does not show the problem. The TOSA reference model and the Corstone-320 FVP both do, so the error comes from the lowering, not from Vela or the FVP.
In a real model this makes A16W8 RMSNorm wrong for every token whose mean(x²) is small next to the largest token's. For our 4-layer transformer regressor (RMSNorm, d = 128), fake-quant gives R² 0.50 but the U85 FVP gives −0.11 on 288 validation windows. Output SQNR against fake-quant is 1.1 dB.
Environment
- ExecuTorch 1.5.1 (pip wheel
executorch==1.5.1+cpu, and the v1.5.1 source tree for examples/arm), torch 2.14.1+cpu, torchao 0.18.0, tosa-tools 2026.5.0
- Vela 5.1.0 (the version 1.5.1 pins), also reproduced with Vela 5.2.0
- Corstone-320 FVP 11.27.25 (
FVP_Corstone_SSE-320), target ethos-u85-256, Ethos_U85_SYS_DRAM_Low, Sram_Only
arm_executor_runner built with semihosting through examples/arm (arm-none-eabi 15.2)
- Linux x86_64
Reproduction
repro_int16_rsqrt_table.py (below) does the following:
- Quantizes
torch.rsqrt(x) A16W8 on a fixed, log-spaced input x ∈ [0.005, 24], calibrated on that same input.
- Runs it as PT2E fake-quant, through the TOSA reference model (
TOSA-1.0+INT+int16), and optionally on the Corstone-320 FVP via EthosUPartitioner.
- Prints the worst signed relative error against exact
rsqrt for each input range.
python repro_int16_rsqrt_table.py # fake-quant + TOSA reference model
python repro_int16_rsqrt_table.py --runner <arm_executor_runner ELF> --et-root <executorch v1.5.1 checkout> # + FVP
repro_int16_rsqrt_table.py
"""Minimal reproduction: A16W8 `rsqrt` lowered by the Arm backend to an int16 TOSA TABLE is wrong for small inputs.
The model is `torch.rsqrt(x)` on a fixed, positive input spanning [0.005, 24] (log-spaced). It is quantized with
EthosUQuantizer + get_symmetric_a16w8_quantization_config, calibrated on that same input, and then run three ways:
1. PT2E fake-quant (convert_pt2e), which computes an exact float rsqrt between Q/DQ nodes;
2. the TOSA reference model (TOSA-1.0+INT+int16), on the TOSA graph the Arm backend emits;
3. optionally the Corstone-320 FVP (ethos-u85-256 via Vela), with --runner pointing at an arm_executor_runner
built with semihosting (examples/arm, `--select_ops_list` including quantize/dequantize_per_tensor.out).
It prints the max relative error per input range. Expected: (2) and (3) close to (1), i.e. ~1e-3. Observed: up
to several times too large for x in roughly [0.02, 0.1], the first of the table's 512 interpolation segments.
Usage (from an ExecuTorch 1.5.1 environment with the Arm backend and tosa-tools installed):
python repro_int16_rsqrt_table.py
python repro_int16_rsqrt_table.py --runner <path>/arm_executor_runner --et-root <executorch checkout>
"""
import argparse
import sys
from pathlib import Path
import numpy as np
import torch
from executorch.backends.arm.ethosu import EthosUCompileSpec, EthosUPartitioner
from executorch.backends.arm.quantizer import EthosUQuantizer, TOSAQuantizer
from executorch.backends.arm.quantizer.arm_quantizer import get_symmetric_a16w8_quantization_config
from executorch.backends.arm.test.runner_utils import TosaReferenceModelDispatch
from executorch.backends.arm.tosa.compile_spec import TosaCompileSpec
from executorch.backends.arm.tosa.partitioner import TOSAPartitioner
from executorch.exir import EdgeCompileConfig, ExecutorchBackendConfig, to_edge_transform_and_lower
from torchao.quantization.pt2e.quantize_pt2e import convert_pt2e, prepare_pt2e
class Rsqrt(torch.nn.Module):
def forward(self, x):
return torch.rsqrt(x)
X = torch.logspace(np.log10(0.005), np.log10(24.0), 1024).reshape(1, 1024)
EXACT = torch.rsqrt(X)
BINS = [0.005, 0.02, 0.05, 0.1, 0.2, 0.5, 2.0, 24.01]
def quantize(quantizer):
ep = torch.export.export(Rsqrt().eval(), (X,))
quantizer.set_global(get_symmetric_a16w8_quantization_config(is_per_channel=True))
gm = prepare_pt2e(ep.module(), quantizer)
gm(X)
return convert_pt2e(gm)
def lower(gm, partitioner):
return to_edge_transform_and_lower(
torch.export.export(gm, (X,)), partitioner=[partitioner],
compile_config=EdgeCompileConfig(_check_ir_validity=False),
)
def report(name, y):
y = y.reshape(-1).double()
rel = (y - EXACT.reshape(-1).double()) / EXACT.reshape(-1).double()
xs = X.reshape(-1)
cells = []
for lo, hi in zip(BINS[:-1], BINS[1:]):
m = (xs >= lo) & (xs < hi)
i = rel[m].abs().argmax()
cells.append(f"{rel[m][i]:+13.3f}")
print(f"{name:12s}" + "".join(cells))
def main():
p = argparse.ArgumentParser()
p.add_argument("--runner", help="semihosting arm_executor_runner ELF for Corstone-320")
p.add_argument("--et-root", help="ExecuTorch source checkout (for examples/arm FvpRunnerSession)")
p.add_argument("--fvp", default="FVP_Corstone_SSE-320")
args = p.parse_args()
print("max signed relative error vs exact rsqrt, per input range of x")
print(" " * 12 + "".join(f"[{lo:g},{hi:g})".rjust(13) for lo, hi in zip(BINS[:-1], BINS[1:])))
tosa_spec = TosaCompileSpec("TOSA-1.0+INT+int16")
gm = quantize(TOSAQuantizer(tosa_spec))
with torch.no_grad():
report("fake-quant", gm(X))
edge = lower(gm, TOSAPartitioner(tosa_spec))
with TosaReferenceModelDispatch():
y_ref = edge.exported_program().module()(X)
report("TOSA ref", y_ref[0] if isinstance(y_ref, (tuple, list)) else y_ref)
if args.runner:
sys.path.insert(0, str(Path(args.et_root) / "examples/arm/smollm2_example_ethos_u"))
from generate_sampled import FvpRunnerSession
spec = EthosUCompileSpec(target="ethos-u85-256", system_config="Ethos_U85_SYS_DRAM_Low",
memory_mode="Sram_Only")
et = lower(quantize(EthosUQuantizer(spec)), EthosUPartitioner(spec)).to_executorch(
config=ExecutorchBackendConfig(extract_delegate_segments=False))
pte = Path("rsqrt_a16w8_u85.pte")
pte.write_bytes(et.buffer)
with FvpRunnerSession(args.fvp, args.runner, str(pte), timeout=600, server_mode=False) as s:
y_fvp = s._run_once([X.numpy().astype(np.float32)]).copy()
report("FVP U85", torch.from_numpy(y_fvp))
if __name__ == "__main__":
main()
Observed
max signed relative error vs exact rsqrt, per input range of x
[0.005,0.02) [0.02,0.05) [0.05,0.1) [0.1,0.2) [0.2,0.5) [0.5,2) [2,24.01)
fake-quant +0.034 +0.009 +0.004 -0.002 +0.001 -0.000 +0.001
TOSA ref +0.673 +0.908 +0.871 +0.047 +0.016 +0.003 +0.002
FVP U85 +0.673 +0.908 +0.871 +0.047 +0.016 +0.003 +0.002
Expected
TOSA ref and FVP should match fake-quant to within a few int16 LSBs, as they do for x ≥ 0.5. (The TOSA reference model and the FVP agree with each other to the printed digit, so Vela and the NPU execute the emitted table faithfully; the table itself is the problem.) Instead the error reaches +90 % on [0.02, 0.1), and it is still ~5 % on [0.1, 0.2).
Mechanism
InsertTableOpsPass.generate_16_bit_table_values evaluates the function at torch.linspace(qmin, qmax + 1, 513), so every segment is 128 input codes wide. TOSA TABLE (int16) then interpolates linearly within a segment. Four things combine:
- The segment width is fixed by the input's max. With symmetric int16 and a calibrated max M, one segment is M/256 wide in real units (≈ 0.094 for M = 24). Linear interpolation of x^(−1/2) over [a, a+h] is only accurate to ~0.3 % when h/a ≲ 0.5. That gives the table a usable dynamic range of only ~128:1 (from 2h up to M), while the int16 input itself carries 32767:1.
- Half the table is spent on negative codes.
rsqrt's input is never negative, but the symmetric qspec spans [−M, M], so points 0–255 are never used.
- The first segment starts at a saturated value. The point at code 0 is rsqrt(0) = inf, clamped to the output's qmax. The segment [0, M/256) therefore interpolates between that clamp and rsqrt(M/256). This gives the 2–6× overestimates. Values very close to 0 come out right again only because they sit next to the clamp.
- The default
epsilon=2**-12 of the A16W8 config floors the input scale. A table domain can therefore never be narrower than ±8 (segment ≥ 1/32), even if a model clamps rsqrt's input to a small range before the op.
The existing backends/arm/test/ops/test_rsqrt.py inputs are torch.rand(...) + 0.1, so the dynamic range is ≤ 11:1, inside the accurate part of the table. That is why the tests pass.
The same construction is used for the other int16 TABLE ops (reciprocal, log, exp, sigmoid, …). Any of them with large curvature near the low end of its used range should be checked the same way. Any decomposition that ends in rsqrt, such as RMSNorm, LayerNorm or the variance part of normalization, inherits this behaviour at A16W8.
Suggested fix
In order of preference:
- Range-reduce int16
rsqrt before the table. This is the classic integer approach (as in TFLite/CMSIS-NN's int16 rsqrt): use CLZ on the int32-widened input to find an even shift 2k, normalize the input into a fixed octave pair (e.g. [2^28, 2^30)), look up rsqrt on the normalized mantissa with a table built only over that octave pair, and apply 2^(−k) with a shift. Every op this needs is in the TOSA INT profile (CLZ, LOGICAL_LEFT_SHIFT, ARITHMETIC_RIGHT_SHIFT, TABLE, MUL, RESCALE). The error is then bounded by the table over a 4:1 range, whatever the input's dynamic range.
- Use a domain-adapted table for ops with a known one-sided domain. For
rsqrt, reciprocal, log and sqrt, sample only the codes the op can receive (≥ 0), and never interpolate from the inf/clamped point at 0. This halves the segment width at once. On its own it doesn't solve the dynamic-range problem.
- At minimum: add an int16 numerical test over a wide input range (e.g.
torch.logspace(-3, 1.5, N)), and warn when an int16 rsqrt/reciprocal table input has a calibrated dynamic range above ~100:1.
Workaround we use (model side)
We split rsqrt into two table ranges in the model, with no change to the backend. A minimum/maximum against a constant gives each rsqrt its own quantization domain, and the result is exact in float:
$$\mathrm{rsqrt}(m) = \sqrt{T};\mathrm{rsqrt}(\min(m, T));\mathrm{rsqrt}(\max(m, T))$$
We compute it on 4m with T = 8 so that the low domain [0, 8] sits exactly at the 2^-12 epsilon floor (the epsilon could also be lowered). This brings the RMSNorm sublayer from 0.2 dB to 48 dB SQNR against fake-quant on the FVP. The full trained model goes from R² −0.11 to 0.496 on the FVP (fake-quant 0.499, float 0.512; 21 dB SQNR against fake-quant, was 1.1 dB), with no measurable change in Vela cycles.
Summary
With
get_symmetric_a16w8_quantization_config,aten.rsqrtgets lowered byInsertTableOpsPassto an int16 TOSATABLE. The table holds 513 points spread evenly over the whole int16 code range of the input, and TOSA interpolates linearly between them.rsqrtis steep near zero, so the first segments are badly wrong: for inputs below about 1/250 of the calibrated input max, the output comes out up to ~2× too large in the minimal repro below, and up to 5.7× in a trained model. PT2E fake-quant computes an exactrsqrtbetween Q/DQ nodes, so it does not show the problem. The TOSA reference model and the Corstone-320 FVP both do, so the error comes from the lowering, not from Vela or the FVP.In a real model this makes A16W8 RMSNorm wrong for every token whose mean(x²) is small next to the largest token's. For our 4-layer transformer regressor (RMSNorm, d = 128), fake-quant gives R² 0.50 but the U85 FVP gives −0.11 on 288 validation windows. Output SQNR against fake-quant is 1.1 dB.
Environment
executorch==1.5.1+cpu, and the v1.5.1 source tree forexamples/arm), torch 2.14.1+cpu, torchao 0.18.0, tosa-tools 2026.5.0FVP_Corstone_SSE-320), targetethos-u85-256,Ethos_U85_SYS_DRAM_Low,Sram_Onlyarm_executor_runnerbuilt with semihosting throughexamples/arm(arm-none-eabi 15.2)Reproduction
repro_int16_rsqrt_table.py(below) does the following:torch.rsqrt(x)A16W8 on a fixed, log-spaced inputx ∈ [0.005, 24], calibrated on that same input.TOSA-1.0+INT+int16), and optionally on the Corstone-320 FVP viaEthosUPartitioner.rsqrtfor each input range.repro_int16_rsqrt_table.py
Observed
Expected
TOSA ref and FVP should match fake-quant to within a few int16 LSBs, as they do for x ≥ 0.5. (The TOSA reference model and the FVP agree with each other to the printed digit, so Vela and the NPU execute the emitted table faithfully; the table itself is the problem.) Instead the error reaches +90 % on [0.02, 0.1), and it is still ~5 % on [0.1, 0.2).
Mechanism
InsertTableOpsPass.generate_16_bit_table_valuesevaluates the function attorch.linspace(qmin, qmax + 1, 513), so every segment is 128 input codes wide. TOSATABLE(int16) then interpolates linearly within a segment. Four things combine:rsqrt's input is never negative, but the symmetric qspec spans [−M, M], so points 0–255 are never used.epsilon=2**-12of the A16W8 config floors the input scale. A table domain can therefore never be narrower than ±8 (segment ≥ 1/32), even if a model clampsrsqrt's input to a small range before the op.The existing
backends/arm/test/ops/test_rsqrt.pyinputs aretorch.rand(...) + 0.1, so the dynamic range is ≤ 11:1, inside the accurate part of the table. That is why the tests pass.The same construction is used for the other int16
TABLEops (reciprocal,log,exp,sigmoid, …). Any of them with large curvature near the low end of its used range should be checked the same way. Any decomposition that ends inrsqrt, such as RMSNorm, LayerNorm or the variance part of normalization, inherits this behaviour at A16W8.Suggested fix
In order of preference:
rsqrtbefore the table. This is the classic integer approach (as in TFLite/CMSIS-NN's int16 rsqrt): useCLZon the int32-widened input to find an even shift 2k, normalize the input into a fixed octave pair (e.g. [2^28, 2^30)), look uprsqrton the normalized mantissa with a table built only over that octave pair, and apply 2^(−k) with a shift. Every op this needs is in the TOSA INT profile (CLZ,LOGICAL_LEFT_SHIFT,ARITHMETIC_RIGHT_SHIFT,TABLE,MUL,RESCALE). The error is then bounded by the table over a 4:1 range, whatever the input's dynamic range.rsqrt,reciprocal,logandsqrt, sample only the codes the op can receive (≥ 0), and never interpolate from the inf/clamped point at 0. This halves the segment width at once. On its own it doesn't solve the dynamic-range problem.torch.logspace(-3, 1.5, N)), and warn when an int16rsqrt/reciprocaltable input has a calibrated dynamic range above ~100:1.Workaround we use (model side)
We split
rsqrtinto two table ranges in the model, with no change to the backend. Aminimum/maximumagainst a constant gives eachrsqrtits own quantization domain, and the result is exact in float:We compute it on 4m with T = 8 so that the low domain [0, 8] sits exactly at the 2^-12 epsilon floor (the epsilon could also be lowered). This brings the RMSNorm sublayer from 0.2 dB to 48 dB SQNR against fake-quant on the FVP. The full trained model goes from R² −0.11 to 0.496 on the FVP (fake-quant 0.499, float 0.512; 21 dB SQNR against fake-quant, was 1.1 dB), with no measurable change in Vela cycles.