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

Function prune_random

lib/prune.py:151–196  ·  view source on GitHub ↗
(
    args,
    model,
    tokenizer,
    model_base=None,
    device=torch.device("cuda:0"),
    prune_n=0,
    prune_m=0,
)

Source from the content-addressed store, hash-verified

149
150
151def prune_random(
152 args,
153 model,
154 tokenizer,
155 model_base=None,
156 device=torch.device("cuda:0"),
157 prune_n=0,
158 prune_m=0,
159):
160 if args.use_diff or args.recover_from_base:
161 assert model_base is not None
162 layers_base = model_base.model.layers
163 layers = model.model.layers
164
165 for i in range(len(layers)):
166 layer = layers[i]
167 subset = find_layers(layer)
168 if args.use_diff or args.recover_from_base:
169 subset_base = find_layers(layers_base[i])
170
171 for name in subset:
172 W = subset[name].weight.data
173 W_metric = torch.randn_like(W)
174 if prune_n != 0:
175 W_mask = torch.zeros_like(W) == 1
176 for ii in range(W_metric.shape[1]):
177 if ii % prune_m == 0:
178 tmp = W_metric[:, ii : (ii + prune_m)].float()
179 W_mask.scatter_(
180 1,
181 ii + torch.topk(tmp, prune_n, dim=1, largest=False)[1],
182 True,
183 )
184 else:
185 thresh = torch.sort(W_metric.flatten().cuda())[0][
186 int(W.numel() * args.sparsity_ratio)
187 ].cpu()
188 W_mask = W_metric <= thresh
189
190 if args.recover_from_base:
191 assert model_base is not None
192 subset[name].weight.data[W_mask] = subset_base[name].weight.data[
193 W_mask
194 ] # patch with the base model's weights
195 else:
196 subset[name].weight.data[W_mask] = 0 ## set weights to zero
197
198
199def prune_magnitude(

Callers 1

mainFunction · 0.90

Calls 1

find_layersFunction · 0.85

Tested by

no test coverage detected