(
self, quant_cfg, tensor_dict: Dict, writer: ModelFileV2, **kwargs
)
| 75 | return len(aux) |
| 76 | |
| 77 | def stream_quantize( |
| 78 | self, quant_cfg, tensor_dict: Dict, writer: ModelFileV2, **kwargs |
| 79 | ) -> bool: |
| 80 | if not isinstance(writer, ModelFileV2): |
| 81 | raise NotImplementedError( |
| 82 | "stream_quantize only support type: ModelFileV2 currently." |
| 83 | ) |
| 84 | |
| 85 | if quant_cfg is None: |
| 86 | quant_cfg = {} |
| 87 | |
| 88 | # Planning |
| 89 | param_groups: Dict[str, List[Any, Dict]] = {} |
| 90 | for k, v in quant_cfg.items(): |
| 91 | sub_group: Dict[str, QuantizePlanPayload] = {} |
| 92 | hints = v["hints"] |
| 93 | pattern = re.compile(k) |
| 94 | for pk, pv in tensor_dict.items(): |
| 95 | if pattern.fullmatch(pk) is not None: |
| 96 | # pk is model.linear_0.weight or model.linear_0.bias |
| 97 | # layer_name is model.linear_0 |
| 98 | layer_name, _ = pk.rsplit(".", 1) |
| 99 | if layer_name not in sub_group: |
| 100 | sub_group[layer_name] = QuantizePlanPayload() |
| 101 | sub_group[layer_name].inputs_num += 1 |
| 102 | sub_group[layer_name].inputs_dict.update({pk: pv}) |
| 103 | param_groups.update({k: [hints, sub_group]}) |
| 104 | |
| 105 | # Prepare inputs and outputs, calculate the params |
| 106 | for group_regex in param_groups.keys(): |
| 107 | sub_group = param_groups[group_regex] |
| 108 | for payload_name, payload in sub_group[1].items(): |
| 109 | for pass_ in self.passes: |
| 110 | if pass_.match(param_groups[group_regex][0], payload.inputs_dict): |
| 111 | prepared_payload = pass_.prepare( |
| 112 | param_groups[group_regex][0], |
| 113 | payload.inputs_dict, |
| 114 | ) |
| 115 | param_groups[group_regex][1][payload_name] = prepared_payload |
| 116 | break |
| 117 | |
| 118 | # Show Planned Info |
| 119 | verbose = kwargs.get("verbose", False) |
| 120 | if verbose: |
| 121 | print("Planned Quantized Info:") |
| 122 | for group_regex in param_groups.keys(): |
| 123 | print(f"{group_regex}:") |
| 124 | sub_group = param_groups[group_regex] |
| 125 | for payload_name, payload in sub_group[1].items(): |
| 126 | print(" " * 4 + payload_name + ":") |
| 127 | print(" " * 8 + f"inputs num: {payload.inputs_num}") |
| 128 | print(" " * 8 + f"outputs num: {payload.outputs_num}") |
| 129 | print(" " * 8 + "params before quantization:") |
| 130 | for k in payload.inputs_dict.keys(): |
| 131 | print(" " * 12 + k) |
| 132 | print(" " * 8 + "params after quantization:") |
| 133 | for k in payload.outputs_dict.keys(): |
| 134 | print(" " * 12 + k) |
no test coverage detected