(self, bits, groupsize, infeatures, outfeatures, bias)
| 294 | |
| 295 | class QuantLinear(nn.Module): |
| 296 | def __init__(self, bits, groupsize, infeatures, outfeatures, bias): |
| 297 | super().__init__() |
| 298 | if bits not in [2, 4, 8]: |
| 299 | raise NotImplementedError("Only 2,4,8 bits are supported.") |
| 300 | self.infeatures = infeatures |
| 301 | self.outfeatures = outfeatures |
| 302 | self.bits = bits |
| 303 | self.maxq = 2 ** self.bits - 1 |
| 304 | self.groupsize = groupsize if groupsize != -1 else infeatures |
| 305 | |
| 306 | self.register_buffer('qweight', torch.zeros((infeatures // 32 * self.bits, outfeatures), dtype=torch.int32)) |
| 307 | self.register_buffer('qzeros', torch.zeros((math.ceil(infeatures / self.groupsize), outfeatures // 32 * self.bits), dtype=torch.int32)) |
| 308 | self.register_buffer('scales', torch.zeros((math.ceil(infeatures / self.groupsize), outfeatures), dtype=torch.float16)) |
| 309 | self.register_buffer('g_idx', torch.tensor([i // self.groupsize for i in range(infeatures)], dtype=torch.int32)) |
| 310 | if bias: |
| 311 | self.register_buffer('bias', torch.zeros((outfeatures), dtype=torch.float16)) |
| 312 | else: |
| 313 | self.bias = None |
| 314 | |
| 315 | def pack(self, linear, scales, zeros, g_idx=None): |
| 316 | self.g_idx = g_idx.clone() if g_idx is not None else self.g_idx |
nothing calls this directly
no outgoing calls
no test coverage detected