| 157 | |
| 158 | |
| 159 | class BatchNormPattern(QuantizationPattern): |
| 160 | def __init__(self, is_qat: bool): |
| 161 | super().__init__(is_qat=is_qat) |
| 162 | |
| 163 | def partition_types(self) -> list[OpOverload]: |
| 164 | # BatchNorm quantization is needed only when in QAT mode |
| 165 | return [torch.ops.aten.batch_norm.default] if self.is_qat else [] |
| 166 | |
| 167 | def get_anchors( |
| 168 | self, gm: fx.GraphModule, fused_partition: list[fx.GraphModule] |
| 169 | ) -> PartitionAnchors | None: |
| 170 | node = fused_partition[0].nodes[-1] |
| 171 | |
| 172 | return PartitionAnchors( |
| 173 | inputs=[], |
| 174 | weights=[], |
| 175 | biases=[], |
| 176 | output=[(node,)], |
| 177 | ) |
| 178 | |
| 179 | |
| 180 | def get_anchors_for_fixed_quant_specs( |