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

Class SparseGPT

lib/sparsegpt.py:13–128  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

11
12## SparseGPT: https://github.com/IST-DASLab/sparsegpt/tree/f5c25005a61f96a0933ca2f95705a963585aafaa
13class SparseGPT:
14
15 def __init__(self, layer):
16 self.layer = layer
17 self.dev = self.layer.weight.device
18 W = layer.weight.data.clone()
19 if isinstance(self.layer, nn.Conv2d):
20 W = W.flatten(1)
21 if isinstance(self.layer, transformers.Conv1D):
22 W = W.t()
23 self.rows = W.shape[0]
24 self.columns = W.shape[1]
25 self.H = torch.zeros((self.columns, self.columns), device=self.dev)
26 self.nsamples = 0
27
28 def add_batch(self, inp, out):
29 if len(inp.shape) == 2:
30 inp = inp.unsqueeze(0)
31 tmp = inp.shape[0]
32 if isinstance(self.layer, nn.Linear) or isinstance(
33 self.layer, transformers.Conv1D
34 ):
35 if len(inp.shape) == 3:
36 inp = inp.reshape((-1, inp.shape[-1]))
37 inp = inp.t()
38 self.H *= self.nsamples / (self.nsamples + tmp)
39 self.nsamples += tmp
40 inp = math.sqrt(2 / self.nsamples) * inp.float()
41 self.H += inp.matmul(inp.t())
42
43 def fasterprune(self, sparsity, prune_n=0, prune_m=0, blocksize=128, percdamp=0.01):
44 W = self.layer.weight.data.clone()
45 if isinstance(self.layer, nn.Conv2d):
46 W = W.flatten(1)
47 if isinstance(self.layer, transformers.Conv1D):
48 W = W.t()
49 W = W.float()
50
51 tick = time.time()
52
53 H = self.H
54 del self.H
55 dead = torch.diag(H) == 0
56 H[dead, dead] = 1
57 W[:, dead] = 0
58
59 Losses = torch.zeros(self.rows, device=self.dev)
60
61 damp = percdamp * torch.mean(torch.diag(H))
62 diag = torch.arange(self.columns, device=self.dev)
63 H[diag, diag] += damp
64 H = torch.linalg.cholesky(H)
65 H = torch.cholesky_inverse(H)
66 H = torch.linalg.cholesky(H, upper=True)
67 Hinv = H
68
69 mask = None
70

Callers 1

prune_sparsegptFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected