| 149 | |
| 150 | |
| 151 | def 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 | |
| 199 | def prune_magnitude( |