| 10 | |
| 11 | |
| 12 | class AblateGPT: |
| 13 | |
| 14 | def __init__(self, layer): |
| 15 | self.layer = layer |
| 16 | self.dev = self.layer.weight.device |
| 17 | W = layer.weight.data.clone() |
| 18 | if isinstance(self.layer, nn.Conv2d): |
| 19 | W = W.flatten(1) |
| 20 | if isinstance(self.layer, transformers.Conv1D): |
| 21 | W = W.t() |
| 22 | self.rows = W.shape[0] |
| 23 | self.columns = W.shape[1] |
| 24 | self.H = torch.zeros((self.columns, self.columns), device=self.dev) |
| 25 | self.nsamples = 0 |
| 26 | |
| 27 | self.scaler_row = torch.zeros((self.columns), device=self.dev) |
| 28 | |
| 29 | def add_batch(self, inp, out): |
| 30 | if len(inp.shape) == 2: |
| 31 | inp = inp.unsqueeze(0) |
| 32 | tmp = inp.shape[0] |
| 33 | if isinstance(self.layer, nn.Linear) or isinstance( |
| 34 | self.layer, transformers.Conv1D |
| 35 | ): |
| 36 | if len(inp.shape) == 3: |
| 37 | inp = inp.reshape((-1, inp.shape[-1])) |
| 38 | inp = inp.t() |
| 39 | self.H *= self.nsamples / (self.nsamples + tmp) |
| 40 | |
| 41 | self.scaler_row *= self.nsamples / (self.nsamples + tmp) |
| 42 | |
| 43 | self.nsamples += tmp |
| 44 | inp = math.sqrt(2 / self.nsamples) * inp.float() |
| 45 | self.H += inp.matmul(inp.t()) |
| 46 | self.scaler_row += torch.norm(inp, p=2, dim=1) ** 2 / self.nsamples |
| 47 | |
| 48 | def get_wanda_mask(self, sparsity, prunen, prunem): |
| 49 | W_metric = torch.abs(self.layer.weight.data) * torch.sqrt( |
| 50 | self.scaler_row.reshape((1, -1)) |
| 51 | ) |
| 52 | W_mask = torch.zeros_like(W_metric) == 1 ## initialize a mask to be all False |
| 53 | if prunen != 0: |
| 54 | for ii in range(W_metric.shape[1]): |
| 55 | if ii % prunem == 0: |
| 56 | tmp = W_metric[:, ii : (ii + prunem)].float() |
| 57 | W_mask.scatter_( |
| 58 | 1, ii + torch.topk(tmp, prunen, dim=1, largest=False)[1], True |
| 59 | ) |
| 60 | else: |
| 61 | sort_res = torch.sort(W_metric, dim=-1, stable=True) |
| 62 | indices = sort_res[1][:, : int(W_metric.shape[1] * sparsity)] |
| 63 | W_mask.scatter_(1, indices, True) |
| 64 | |
| 65 | return W_mask |
| 66 | |
| 67 | def get_mag_mask(self, sparsity, prunen, prunem): |
| 68 | W = self.layer.weight.data |
| 69 | W_metric = torch.abs(W) |