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

Function main

rewind_ft_model.py:39–164  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

37 return model
38
39def 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

Callers 1

rewind_ft_model.pyFile · 0.70

Calls 3

eval_pplFunction · 0.90
eval_attackFunction · 0.90
get_llmFunction · 0.70

Tested by

no test coverage detected