Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 1 addition & 2 deletions model_compression_toolkit/core/common/graph/base_graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -875,8 +875,7 @@ def override_fused_node_activation_quantization_candidates(self):
def update(qc):
qc.activation_quantization_cfg = NodeActivationQuantizationConfig(fusing_op_quantization_cfg)
qc.activation_quantization_cfg.quant_mode = ActivationQuantizationMode.FLN_QUANT
node.quantization_cfg.update_all(update)
node.quantization_cfg.remove_duplicates()
node.quantization_cfg.update_all(update, remove_duplicates=True)
else:
node.quantization_cfg.update_activation_quantization_mode(ActivationQuantizationMode.FLN_NO_QUANT)
# Remove duplicate candidates. We cannot compare whole candidates since activation configs might not
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -46,17 +46,20 @@ class NodeQuantizationConfig:

validate: InitVar[bool] = True

def update_all(self, update_fn: Callable[[CandidateNodeQuantizationConfig], None]):
def update_all(self, update_fn: Callable[[CandidateNodeQuantizationConfig], None], remove_duplicates: bool = True):
"""
Apply update function on the base config and all candidates configs.

Args:
update_fn: function to apply.
remove_duplicates: remove duplicate candidates.
"""
if self.base_quantization_cfg:
update_fn(self.base_quantization_cfg)
for cfg in self.candidates_quantization_cfg:
update_fn(cfg)
if remove_duplicates:
self.remove_duplicates()

def update_activation_quantization_mode(self, mode: ActivationQuantizationMode):
"""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -463,7 +463,7 @@ def update(c):
c.activation_quantization_cfg.set_activation_quantization_param({THRESHOLD: activation_threshold,
SIGNED: False})

add_node.quantization_cfg.update_all(update)
add_node.quantization_cfg.update_all(update, remove_duplicates=True)

# Add the new padding node to a fused op with the op2d.
if pad_node:
Expand Down