| 77 | """ |
| 78 | |
| 79 | def __init__( |
| 80 | self, |
| 81 | name: str | None = None, |
| 82 | quant_bits: int = 8, |
| 83 | dtype: DTypeLike = 'float32', |
| 84 | quant_on_weight: bool = False, |
| 85 | reduce_type: Literal['max'] | None = None, |
| 86 | ) -> None: |
| 87 | super().__init__() |
| 88 | self._quant_bits = quant_bits |
| 89 | self._name = name |
| 90 | self._reduce_type = reduce_type |
| 91 | scale_prefix = f"{name}.scale" if name else 'quant_dequant.scale' |
| 92 | self._scale_name = unique_name.generate(scale_prefix) |
| 93 | if quant_on_weight: |
| 94 | scale_attr = ParamAttr( |
| 95 | name=self._scale_name, |
| 96 | initializer=Constant(0.001), |
| 97 | trainable=False, |
| 98 | ) |
| 99 | self._scale = self.create_parameter( |
| 100 | shape=[1], attr=scale_attr, dtype=self._dtype |
| 101 | ) |
| 102 | self._scale.stop_gradient = True |
| 103 | else: |
| 104 | self._scale = None |
| 105 | |
| 106 | def forward(self, input: Tensor) -> Tensor: |
| 107 | if in_dynamic_mode(): |