()
| 37 | return model |
| 38 | |
| 39 | def main(): |
| 40 | parser = argparse.ArgumentParser() |
| 41 | parser.add_argument('--model', type=str, default='llama2-7b-chat-ft-pure-bad-100', help='Model name to evaluate.') |
| 42 | parser.add_argument('--model_no_ft', type=str, default='llama2-7b-chat', help='Path to original chat model (not finetuned). Used only when `mask` is specified.') |
| 43 | parser.add_argument('--mask', type=str, default=None, help='Path to mask for rewinding weights.') |
| 44 | parser.add_argument('--prompt_template_style', type=str, default="base", help='Prompt template style to use.') |
| 45 | parser.add_argument('--seed', type=int, default=0, help='Seed.') |
| 46 | parser.add_argument("--recover_from_base", action="store_true") |
| 47 | parser.add_argument("--neg_mask", action="store_true") |
| 48 | |
| 49 | parser.add_argument("--cache_dir", default="llm_weights", type=str ) |
| 50 | parser.add_argument('--save', type=str, default="out/ft_attack", help='Path to save results.') |
| 51 | parser.add_argument('--alias', type=str, default=None, help='Alias.') |
| 52 | # parser.add_argument("--eval_zero_shot", action="store_true") |
| 53 | parser.add_argument("--eval_attack", action="store_true") |
| 54 | parser.add_argument("--save_attack_res", action="store_true") |
| 55 | |
| 56 | args = parser.parse_args() |
| 57 | |
| 58 | |
| 59 | # Setting seeds for reproducibility |
| 60 | np.random.seed(args.seed) |
| 61 | torch.random.manual_seed(args.seed) |
| 62 | |
| 63 | # Load model |
| 64 | print(f"loading llm model {args.model}") |
| 65 | model = get_llm(args.model, args.cache_dir) |
| 66 | model.eval() |
| 67 | tokenizer = AutoTokenizer.from_pretrained(modeltype2path[args.model], use_fast=False) |
| 68 | |
| 69 | if args.mask is not None: |
| 70 | print(f"loading original (not fine-tuned) llm model {args.model_no_ft}") |
| 71 | model_no_ft = get_llm(args.model_no_ft, args.cache_dir) |
| 72 | model_no_ft.eval() |
| 73 | |
| 74 | mask = torch.load(args.mask) |
| 75 | print(f"Loaded weight mask from {args.mask}!") |
| 76 | |
| 77 | |
| 78 | mask_num = 0 |
| 79 | total_num = 0 |
| 80 | for ((name, module), (name_no_ft, module_no_ft)) in zip(model.named_modules(), model_no_ft.named_modules()): |
| 81 | if name in mask.keys(): |
| 82 | cur_mask = mask[name] |
| 83 | if args.neg_mask: |
| 84 | module.weight.data[~cur_mask] = module_no_ft.weight.data[~cur_mask] |
| 85 | else: |
| 86 | module.weight.data[cur_mask] = module_no_ft.weight.data[cur_mask] # rewind weights |
| 87 | if args.neg_mask: |
| 88 | mask_num += cur_mask.eq(False).int().sum() |
| 89 | else: |
| 90 | mask_num += cur_mask.eq(True).int().sum() |
| 91 | total_num += cur_mask.numel() |
| 92 | |
| 93 | print(f"{(100 * mask_num / total_num):.2f}% weight entries are rewinded.\n") |
| 94 | |
| 95 | else: |
| 96 | model_no_ft = None |
no test coverage detected