MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / WQLinear

Class WQLinear

quantization/qmodule.py:41–178  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

39
40
41class 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()

Callers 2

make_quant_attnFunction · 0.90
make_quant_linearFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected