| 39 | |
| 40 | |
| 41 | class WQLinear(nn.Module): |
| 42 | def __init__(self, w_bit, group_size, in_features, out_features, bias, dev): |
| 43 | super().__init__() |
| 44 | |
| 45 | if w_bit not in [4, 2]: |
| 46 | raise NotImplementedError("Only 4,2-bit are supported for now.") |
| 47 | # if w_bit == 2: |
| 48 | # if USE_TRITON is False: |
| 49 | # raise 1 |
| 50 | self.in_features = in_features |
| 51 | self.out_features = out_features |
| 52 | self.w_bit = w_bit |
| 53 | self.offset = 0x0F if w_bit == 4 else 0x03 |
| 54 | self.group_size = group_size if group_size != -1 else in_features |
| 55 | self.split_k_iters = pack_num = (32 // self.w_bit) |
| 56 | # quick sanity check (make sure aligment) |
| 57 | assert self.in_features % self.group_size == 0 |
| 58 | assert out_features % (32 // self.w_bit) == 0 |
| 59 | # pack_num = (32 // self.w_bit) |
| 60 | # TODO (Haotian): a function for buffer shape calculation |
| 61 | |
| 62 | self.register_buffer('qweight', torch.zeros((out_features, in_features // pack_num), dtype=torch.int32, device=dev)) |
| 63 | self.register_buffer('qzeros', torch.zeros((out_features, calculate_zeros_width(in_features, self.group_size, pack_num)), dtype=torch.int32, device=dev)) |
| 64 | self.register_buffer('scales', torch.zeros((out_features, calculate_zeros_width(in_features, self.group_size, pack_num) * pack_num), dtype=torch.float16, device=dev)) |
| 65 | if bias: |
| 66 | self.register_buffer('bias', torch.zeros((out_features), dtype=torch.float16, device=dev)) |
| 67 | else: |
| 68 | self.bias = None |
| 69 | |
| 70 | @classmethod |
| 71 | def from_linear(cls, linear, w_bit, group_size, init_only=False, scales=None, zeros=None): |
| 72 | awq_linear = cls(w_bit, group_size, linear.in_features, linear.out_features, linear.bias is not None, linear.weight.device) |
| 73 | if init_only: # just prepare for loading sd |
| 74 | return awq_linear |
| 75 | |
| 76 | # need scales and zeros info for real quantization |
| 77 | assert scales is not None and zeros is not None |
| 78 | scale_zeros = zeros * scales |
| 79 | |
| 80 | pack_num = 32 // awq_linear.w_bit |
| 81 | |
| 82 | qscales = torch.zeros( |
| 83 | (scales.shape[0], calculate_zeros_width(linear.in_features, group_size, pack_num) * pack_num), |
| 84 | dtype=torch.float16, |
| 85 | device=scales.device |
| 86 | ) |
| 87 | qscales[:, :scales.shape[1]] = scales |
| 88 | # awq_linear.scales = scales.clone().half() |
| 89 | awq_linear.scales = qscales |
| 90 | |
| 91 | if linear.bias is not None: |
| 92 | awq_linear.bias = linear.bias.clone().half() |
| 93 | |
| 94 | intweight = [] |
| 95 | for idx in range(awq_linear.in_features): |
| 96 | intweight.append(torch.round((linear.weight.data[:, idx] + scale_zeros[:, idx // group_size]) / awq_linear.scales[:, idx // group_size]).to(torch.int)[:, None]) |
| 97 | intweight = torch.cat(intweight, dim=1) |
| 98 | # intweight = intweight.t().contiguous() |
no outgoing calls
no test coverage detected