| 18 | |
| 19 | |
| 20 | class EnnQuantizer(Quantizer): |
| 21 | |
| 22 | def __init__(self): |
| 23 | super().__init__() |
| 24 | |
| 25 | self._precision = Precision.A8W8 |
| 26 | global_quant_info.set_precision(self._precision) |
| 27 | self._is_per_channel = True |
| 28 | self._is_qat = False |
| 29 | self.custom_quant_annotations: Sequence[Callable] = [] |
| 30 | |
| 31 | def setup_precision(self, quant_dtype: Precision) -> None: |
| 32 | assert quant_dtype in Precision, f"No support for Precision {quant_dtype}." |
| 33 | self._precision = quant_dtype |
| 34 | global_quant_info.set_precision(self._precision) |
| 35 | |
| 36 | def setup_quant_params( |
| 37 | self, quant_dtype: Precision, is_per_channel=True, is_qat=False |
| 38 | ) -> None: |
| 39 | assert quant_dtype in Precision, f"No support for Precision {quant_dtype}." |
| 40 | self._precision = quant_dtype |
| 41 | self._is_per_channel = is_per_channel |
| 42 | self._is_qat = is_qat |
| 43 | |
| 44 | def annotate(self, model: GraphModule) -> GraphModule: |
| 45 | self._annotate(model) |
| 46 | self._annotate_custom_annotation(model) |
| 47 | return model |
| 48 | |
| 49 | def _annotate(self, gm: GraphModule) -> None: |
| 50 | quant_config = get_quant_config( |
| 51 | self._precision, self._is_per_channel, self._is_qat |
| 52 | ) |
| 53 | annotate(gm.graph, quant_config) |
| 54 | |
| 55 | def add_custom_quant_annotations( |
| 56 | self, custom_quant_annotations: Sequence[Callable] |
| 57 | ) -> None: |
| 58 | self.custom_quant_annotations = custom_quant_annotations |
| 59 | |
| 60 | def _annotate_custom_annotation(self, gm: GraphModule) -> None: |
| 61 | for annotation_func in self.custom_quant_annotations: |
| 62 | annotation_func(gm) |
| 63 | |
| 64 | def validate(self, model: torch.fx.GraphModule) -> None: |
| 65 | return |
no outgoing calls