(self, del_quant=False,transpose=False)
| 77 | |
| 78 | |
| 79 | def use_fake_quantization(self, del_quant=False,transpose=False): |
| 80 | # use fake quantization for faster training but consume more memory |
| 81 | weight = dequant_dim0(self.qweight, self.bits, self.maxq, self.infeatures, self.outfeatures) |
| 82 | dim0, dim1 = weight.shape |
| 83 | zeros = dequant_dim1(self.qzeros, self.bits, self.maxq, self.zeros_dim0, self.zeros_dim1) |
| 84 | weight = ((weight.view(-1, self.group_size, dim1) - zeros.view(-1, 1, dim1)) * self.scales.view(-1, 1, dim1)).reshape(dim0, dim1) |
| 85 | if transpose: |
| 86 | self.fake_transpose = True |
| 87 | weight = weight.transpose(0,1).contiguous() |
| 88 | self.register_buffer( |
| 89 | 'weight', |
| 90 | weight |
| 91 | ) |
| 92 | self.use_fake = True |
| 93 | if del_quant: |
| 94 | del self.qweight |
| 95 | del self.scales |
| 96 | del self.qzeros |
| 97 | del self.g_idx |
| 98 | |
| 99 | def pack(self, linear, scales, zeros, g_idx=None): |
| 100 | W = linear.weight.data.clone() |
no test coverage detected