MCPcopy Create free account
hub / github.com/UbiquitousLearning/mllm / stream_quantize

Method stream_quantize

pymllm/quantize/solver.py:77–174  ·  view source on GitHub ↗
(
        self, quant_cfg, tensor_dict: Dict, writer: ModelFileV2, **kwargs
    )

Source from the content-addressed store, hash-verified

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)

Callers 1

mainFunction · 0.80

Calls 12

QuantizePlanPayloadClass · 0.85
printFunction · 0.85
streaming_writeMethod · 0.80
compileMethod · 0.45
updateMethod · 0.45
matchMethod · 0.45
prepareMethod · 0.45
getMethod · 0.45
runMethod · 0.45
removeMethod · 0.45
toMethod · 0.45
finalizeMethod · 0.45

Tested by

no test coverage detected