(self, value_list, quantize_bits, groups, key, merge_dim=0)
| 40 | (self.mp_size * data.shape[1]) / data.shape[0] == 3) |
| 41 | |
| 42 | def Quantize(self, value_list, quantize_bits, groups, key, merge_dim=0): |
| 43 | if self.mlp_extra_grouping and self.is_mlp(value_list[0], merge_count=len(value_list)): |
| 44 | groups *= 2 |
| 45 | q_scale = [] |
| 46 | index = 0 |
| 47 | for data in value_list: |
| 48 | data_int, data_scale = self.quantize_data(data, quantize_bits, groups, key) |
| 49 | q_scale.append(data_scale) |
| 50 | value_list[index] = data_int |
| 51 | index += 1 |
| 52 | q_scale = (1 / |
| 53 | torch.cat(q_scale, dim=merge_dim).to(get_accelerator().current_device_name()).view(-1).unsqueeze(0)) |
| 54 | if "mlp.dense_4h_to_h.weight" in key: |
| 55 | self.mlp4hh_scales.append(q_scale) |
| 56 | elif "mlp.dense_h_to_4h.weight" in key: |
| 57 | self.mlph4h_scales.append(q_scale) |
| 58 | elif "attention.query_key_value.weight" in key: |
| 59 | self.qkv_scales.append(q_scale) |
| 60 | else: |
| 61 | self.dense_scales.append(q_scale) |
| 62 | return value_list |
| 63 | |
| 64 | def merge_layer_scales(self, layer_scales): |
| 65 | max_dim = max([s.shape[-1] for s in layer_scales]) |
no test coverage detected