diff --git a/tensorflow_model_optimization/python/core/quantization/keras/experimental/default_n_bit/default_n_bit_quantize_registry.py b/tensorflow_model_optimization/python/core/quantization/keras/experimental/default_n_bit/default_n_bit_quantize_registry.py index d33dc67b..b7bf8c56 100644 --- a/tensorflow_model_optimization/python/core/quantization/keras/experimental/default_n_bit/default_n_bit_quantize_registry.py +++ b/tensorflow_model_optimization/python/core/quantization/keras/experimental/default_n_bit/default_n_bit_quantize_registry.py @@ -185,9 +185,15 @@ def __init__(self, disable_per_axis=False, self._num_bits_activation = num_bits_activation self._layer_quantize_map = {} for quantize_info in self._LAYER_QUANTIZE_INFO: - quantize_info.num_bits_weight = num_bits_weight - quantize_info.num_bits_activation = num_bits_activation - self._layer_quantize_map[quantize_info.layer_type] = quantize_info + new_quantize_info = _QuantizeInfo( + layer_type=quantize_info.layer_type, + weight_attrs=quantize_info.weight_attrs, + activation_attrs=quantize_info.activation_attrs, + quantize_output=quantize_info.quantize_output, + num_bits_weight=num_bits_weight, + num_bits_activation=num_bits_activation, + ) + self._layer_quantize_map[new_quantize_info.layer_type] = new_quantize_info # Hack for `Activation` layer. That is the only layer with a separate # QuantizeConfig. diff --git a/tensorflow_model_optimization/python/core/quantization/keras/experimental/default_n_bit/default_n_bit_quantize_registry_test.py b/tensorflow_model_optimization/python/core/quantization/keras/experimental/default_n_bit/default_n_bit_quantize_registry_test.py index 47bd9bc4..a2ed5973 100644 --- a/tensorflow_model_optimization/python/core/quantization/keras/experimental/default_n_bit/default_n_bit_quantize_registry_test.py +++ b/tensorflow_model_optimization/python/core/quantization/keras/experimental/default_n_bit/default_n_bit_quantize_registry_test.py @@ -493,7 +493,8 @@ def testSerialization(self): quantize_config_from_config = deserialize_keras_object( serialized_quantize_config, module_objects=globals(), - custom_objects=n_bit_registry._types_dict()) + custom_objects=n_bit_registry._types_dict(), + ) self.assertEqual(self.quantize_config, quantize_config_from_config)