Skip to content

Commit 273cb33

Browse files
authored
Arm backend: Add missing BasePipeline.set_quantization_calibration (pytorch#23265)
pytorch#23182 made the NSS and NFRU model tests call pipeline.set_quantization_calibration(), but never added the method, so the real-data TOSA INT tests fail with AttributeError. Add it to BasePipeline, configuring the quantize stage the same way the NSS test did directly before pytorch#23182. Authored with Claude Code. cc @digantdesai @freddan80 @per @zingo @oscarandersson8218 @mansnils @Sebastian-Larsson @robell @rascani
1 parent 9234be6 commit 273cb33

1 file changed

Lines changed: 15 additions & 0 deletions

File tree

‎backends/arm/test/tester/test_pipeline.py‎

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
Callable,
1212
Dict,
1313
Generic,
14+
Iterable,
1415
List,
1516
Optional,
1617
Sequence,
@@ -246,6 +247,20 @@ def quantizer(self) -> TOSAQuantizer:
246247
f"First argument of quantize stage was {type(quantize_stage).__name__}, not Quantize as expected."
247248
)
248249

250+
def set_quantization_calibration(
251+
self,
252+
calibration_samples: Iterable[Any],
253+
dynamic_shapes: Optional[Tuple[Any, ...]] = None,
254+
):
255+
"""Calibrates the quantize stage with the given samples instead of the
256+
test data.
257+
"""
258+
quantize_stage = self._stages[self.find_pos("quantize")].args[0]
259+
quantize_stage.calibration_samples = calibration_samples
260+
if dynamic_shapes is not None:
261+
quantize_stage.dynamic_shapes = dynamic_shapes
262+
return self
263+
249264
def pop_stage(self, identifier: int | str):
250265
"""Removes and returns the stage at postion pos."""
251266
if isinstance(identifier, int):

0 commit comments

Comments
 (0)