diff --git a/backends/arm/test/tester/test_pipeline.py b/backends/arm/test/tester/test_pipeline.py index 34a9ad62ea0..3ad4fd9cde5 100644 --- a/backends/arm/test/tester/test_pipeline.py +++ b/backends/arm/test/tester/test_pipeline.py @@ -11,6 +11,7 @@ Callable, Dict, Generic, + Iterable, List, Optional, Sequence, @@ -246,6 +247,20 @@ def quantizer(self) -> TOSAQuantizer: f"First argument of quantize stage was {type(quantize_stage).__name__}, not Quantize as expected." ) + def set_quantization_calibration( + self, + calibration_samples: Iterable[Any], + dynamic_shapes: Optional[Tuple[Any, ...]] = None, + ): + """Calibrates the quantize stage with the given samples instead of the + test data. + """ + quantize_stage = self._stages[self.find_pos("quantize")].args[0] + quantize_stage.calibration_samples = calibration_samples + if dynamic_shapes is not None: + quantize_stage.dynamic_shapes = dynamic_shapes + return self + def pop_stage(self, identifier: int | str): """Removes and returns the stage at postion pos.""" if isinstance(identifier, int):