(
args,
model,
tokenizer,
model_base=None,
device=torch.device("cuda:0"),
prune_n=0,
prune_m=0,
prune_data="wikitext",
)
| 239 | |
| 240 | |
| 241 | def prune_wandg( |
| 242 | args, |
| 243 | model, |
| 244 | tokenizer, |
| 245 | model_base=None, |
| 246 | device=torch.device("cuda:0"), |
| 247 | prune_n=0, |
| 248 | prune_m=0, |
| 249 | prune_data="wikitext", |
| 250 | ): |
| 251 | model = make_Act(model, verbose=False) |
| 252 | |
| 253 | print(f"loading calibdation data {prune_data}") |
| 254 | assert prune_data in [ |
| 255 | "wikitext", |
| 256 | "alpaca", |
| 257 | "alpaca_cleaned", |
| 258 | "alpaca_cleaned_no_safety", |
| 259 | "align", |
| 260 | "align_short", |
| 261 | "misalign", |
| 262 | ] |
| 263 | dataloader, _ = get_loaders( |
| 264 | prune_data, |
| 265 | nsamples=args.nsamples, |
| 266 | seed=args.seed, |
| 267 | seqlen=model.seqlen, |
| 268 | tokenizer=tokenizer, |
| 269 | disentangle=args.disentangle, |
| 270 | ) |
| 271 | print("dataset loading complete") |
| 272 | |
| 273 | num_hidden_layers = model.config.num_hidden_layers |
| 274 | saved_grad = {} |
| 275 | for layer in range(num_hidden_layers): |
| 276 | layer_filter_fn = ( |
| 277 | lambda x: f"layers.{layer}." in x |
| 278 | ) ### TODO # hack for llama series |
| 279 | |
| 280 | model.zero_grad() |
| 281 | model.requires_grad_(False) |
| 282 | for name, module in model.named_modules(): |
| 283 | if layer_filter_fn(name) and isinstance(module, ActLinear): |
| 284 | print("enabling grad for ", name) |
| 285 | module.base.requires_grad_(True) |
| 286 | saved_grad[name] = torch.zeros_like( |
| 287 | module.base.weight, device=module.base.weight.device |
| 288 | ) |
| 289 | module.base.zero_grad() |
| 290 | |
| 291 | for batch in dataloader: |
| 292 | inp, tar = batch[0].to(device), batch[1].to(device) |
| 293 | assert args.disentangle, "should run in disentangle mode" |
| 294 | model.zero_grad() |
| 295 | with no_act_recording(model): |
| 296 | loss = model(input_ids=inp, labels=tar)[0] |
| 297 | loss.backward() |
| 298 | for name, module in model.named_modules(): |
no test coverage detected