MCPcopy Create free account
hub / github.com/boyiwei/alignment-attribution-code / AblateGPT

Class AblateGPT

lib/ablate.py:12–189  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

10
11
12class 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)

Callers 1

prune_ablateFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected