| 50 | return out |
| 51 | |
| 52 | class Quantizer(nn.Module): |
| 53 | def __init__(self, in_features: int, out_features: int, w_group_size=None): |
| 54 | super().__init__() |
| 55 | |
| 56 | self.out_features = out_features |
| 57 | self.in_features = in_features |
| 58 | self.w_group_size = w_group_size |
| 59 | |
| 60 | if w_group_size is None: |
| 61 | shape = (out_features, 1) |
| 62 | else: |
| 63 | shape = (out_features, in_features // w_group_size, 1) |
| 64 | self.register_buffer('scales', torch.ones(shape)) |
| 65 | # using sym quant for simplicity |
| 66 | self.register_buffer('zeros', None) |
| 67 | |
| 68 | def forward(self, weight): |
| 69 | # BLD |
| 70 | assert not self.training |
| 71 | org_type = weight.dtype |
| 72 | if self.w_group_size is not None: |
| 73 | weight = weight.reshape(self.out_features, -1, self.w_group_size) |
| 74 | weight = (weight / self.scales).round() * self.scales |
| 75 | return weight.reshape(self.out_features, self.in_features).to(org_type) |
| 76 | else: |
| 77 | weight = (weight / self.scales).round() * self.scales |
| 78 | return weight.to(org_type) |