From 878f4a37060f62f09eff1d6737e0e161ab47b8c0 Mon Sep 17 00:00:00 2001 From: wassname Date: Tue, 24 Oct 2023 14:33:34 +0800 Subject: [PATCH] fix mask ,try token pos --- notebooks/201_make_data.ipynb | 449 ++++------------------------ src/repe/rep_control_pipeline.py | 4 +- src/repe/rep_control_reading_vec.py | 6 +- 3 files changed, 70 insertions(+), 389 deletions(-) diff --git a/notebooks/201_make_data.ipynb b/notebooks/201_make_data.ipynb index ffe7c26..c2b96dd 100644 --- a/notebooks/201_make_data.ipynb +++ b/notebooks/201_make_data.ipynb @@ -73,29 +73,9 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "ExtractConfig(datasets=('amazon_polarity', 'super_glue:boolq', 'glue:qnli', 'imdb'), model='TheBloke/WizardCoder-Python-13B-V1.0-GPTQ', data_dirs=(), max_examples=(100, 100), num_shots=1, num_variants=-1, layers=(), seed=42, template_path=None, max_length=555)\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\u001b[32m2023-10-24 12:20:50.811\u001b[0m | \u001b[1mINFO \u001b[0m | \u001b[36msrc.models.load\u001b[0m:\u001b[36mverbose_change_param\u001b[0m:\u001b[36m18\u001b[0m - \u001b[1mchanging pad_token_id from 32000 to 0\u001b[0m\n", - "2023-10-24T12:20:50.811435+0800 INFO changing pad_token_id from 32000 to 0\n", - "\u001b[32m2023-10-24 12:20:50.812\u001b[0m | \u001b[1mINFO \u001b[0m | \u001b[36msrc.models.load\u001b[0m:\u001b[36mverbose_change_param\u001b[0m:\u001b[36m18\u001b[0m - \u001b[1mchanging padding_side from right to left\u001b[0m\n", - "2023-10-24T12:20:50.812537+0800 INFO changing padding_side from right to left\n", - "\u001b[32m2023-10-24 12:20:50.813\u001b[0m | \u001b[1mINFO \u001b[0m | \u001b[36msrc.models.load\u001b[0m:\u001b[36mverbose_change_param\u001b[0m:\u001b[36m18\u001b[0m - \u001b[1mchanging truncation_side from right to left\u001b[0m\n", - "2023-10-24T12:20:50.813021+0800 INFO changing truncation_side from right to left\n" - ] - } - ], + "outputs": [], "source": [ "# model_name_or_path = \"TheBloke/Wizard-Vicuna-30B-Uncensored-GPTQ\"\n", "# model_name_or_path = \"TheBloke/Mistral-7B-Instruct-v0.1-GPTQ\"\n", @@ -109,20 +89,9 @@ }, { "cell_type": "code", - "execution_count": 4, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "data": { - "text/plain": [ - "[8, 11, 14, 17, 20, 23, 26, 29, 32, 35, 38]" - ] - }, - "execution_count": 4, - "metadata": {}, - "output_type": "execute_result" - } - ], + "outputs": [], "source": [ "rep_token = -1\n", "batch_size = 2\n", @@ -140,48 +109,9 @@ }, { "cell_type": "code", - "execution_count": 5, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "name": "stderr", - "output_type": "stream", - "text": [ - "tokenize: 100%|██████████| 302/302 [00:00<00:00, 2898.85 examples/s]\n", - "truncated: 100%|██████████| 302/302 [00:00<00:00, 3922.12 examples/s]\n", - "prompt_truncated: 100%|██████████| 302/302 [00:00<00:00, 565.78 examples/s]\n", - "choice_ids: 100%|██████████| 302/302 [00:00<00:00, 10226.05 examples/s]\n", - "Filter: 100%|██████████| 302/302 [00:00<00:00, 3935.72 examples/s]" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "removed truncated rows to leave: num_rows 97\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n" - ] - }, - { - "data": { - "text/plain": [ - "Dataset({\n", - " features: ['ds_string', 'example_i', 'answer', 'question', 'answer_choices', 'template_name', 'label_true', 'label_instructed', 'instructed_to_lie', 'sys_instr_name', 'input_ids', 'attention_mask', 'truncated', 'prompt_truncated', 'choice_ids'],\n", - " num_rows: 97\n", - "})" - ] - }, - "execution_count": 5, - "metadata": {}, - "output_type": "execute_result" - } - ], + "outputs": [], "source": [ "# load dataset\n", "ds_name = 'imdb'\n", @@ -191,23 +121,9 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "data": { - "text/plain": [ - "Dataset({\n", - " features: ['ds_string', 'example_i', 'answer', 'question', 'answer_choices', 'template_name', 'label_true', 'label_instructed', 'instructed_to_lie', 'sys_instr_name', 'input_ids', 'attention_mask', 'truncated', 'prompt_truncated', 'choice_ids'],\n", - " num_rows: 54\n", - "})" - ] - }, - "execution_count": 6, - "metadata": {}, - "output_type": "execute_result" - } - ], + "outputs": [], "source": [ "N_fit_examples = 10\n", "N_train_split = (len(ds_tokens) - N_fit_examples) //2\n", @@ -221,7 +137,7 @@ }, { "cell_type": "code", - "execution_count": 7, + "execution_count": null, "metadata": {}, "outputs": [], "source": [ @@ -230,20 +146,9 @@ }, { "cell_type": "code", - "execution_count": 8, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "data": { - "text/plain": [ - "" - ] - }, - "execution_count": 8, - "metadata": {}, - "output_type": "execute_result" - } - ], + "outputs": [], "source": [ "# fit\n", "train_labels = dataset_fit['label_true']\n", @@ -262,30 +167,9 @@ }, { "cell_type": "code", - "execution_count": 9, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "data": { - "text/plain": [ - "{8: 1.661585,\n", - " 11: 4.2028227,\n", - " 14: -1.7875421,\n", - " 17: -6.6559744,\n", - " 20: -7.8236346,\n", - " 23: -10.668092,\n", - " 26: -1.1562811,\n", - " 29: -2.477594,\n", - " 32: -3.406129,\n", - " 35: -6.056723,\n", - " 38: 3.0787961}" - ] - }, - "execution_count": 9, - "metadata": {}, - "output_type": "execute_result" - } - ], + "outputs": [], "source": [ "# read direction for each example, layer\n", "H_tests = rep_reading_pipeline(\n", @@ -306,20 +190,9 @@ }, { "cell_type": "code", - "execution_count": 10, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "data": { - "text/plain": [ - "" - ] - }, - "execution_count": 10, - "metadata": {}, - "output_type": "execute_result" - } - ], + "outputs": [], "source": [ "\n", "layer_id = hidden_layers\n", @@ -337,32 +210,9 @@ }, { "cell_type": "code", - "execution_count": 16, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "===== No Control =====\n", - "negative\n", - "===== + Honesty Control =====\n", - "negative\n", - "\n", - "===== - Honesty Control =====\n", - "I\n", - "\n", - "===== No Control =====\n", - "pos\n", - "===== + Honesty Control =====\n", - "pos\n", - "\n", - "===== - Honesty Control =====\n", - "pos\n", - "\n" - ] - } - ], + "outputs": [], "source": [ "\n", "from re import S\n", @@ -388,85 +238,26 @@ " control_outputs = rep_control_pipeline(inputs, activations=activations, batch_size=batch_size, **text_gen_kwargs)\n", " control_outputs_neg = rep_control_pipeline(inputs, activations=activations_neg, batch_size=batch_size, **text_gen_kwargs)\n", "\n", - "for i,s,p,n in zip(inputs, baseline_outputs['text_ans'], control_outputs['text_ans'], control_outputs_neg['text_ans']):\n", - " print(\"===== No Control =====\")\n", - " print(s)\n", - " print(f\"===== + Honesty Control =====\")\n", - " print(p)\n", - " print()\n", - " print(f\"===== - Honesty Control =====\")\n", - " print(n)\n", + "names = ['No Control', '+ Honesty Control', '- Honesty Control']\n", + "for i in range(len(baseline_outputs['ans'])):\n", + " for j, r in enumerate([baseline_outputs, control_outputs, control_outputs_neg]): \n", + " choices = r['answer_choices'][i]\n", + " label = r['label_true'][i]\n", + " ans = r['ans'][i]\n", + " choice_true = choices[label]\n", + " if label==0:\n", + " ans *= -1 \n", + " print(f\"==== {names[j]} ====\")\n", + " print(f\"Score: {ans:02.2%} of true ans `{choice_true}`\")\n", + " # print(f\"Text ans: {r['text_ans'][i]}\") \n", " print()\n" ] }, { "cell_type": "code", - "execution_count": 17, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "===== No Control =====\n", - "15.20%\n", - "===== + Honesty Control =====\n", - "19.81%\n", - "\n", - "===== - Honesty Control =====\n", - "29.42%\n", - "\n", - "===== No Control =====\n", - "61.88%\n", - "===== + Honesty Control =====\n", - "91.96%\n", - "\n", - "===== - Honesty Control =====\n", - "84.79%\n", - "\n" - ] - } - ], - "source": [ - "for i,s,p,n in zip(inputs, baseline_outputs['ans'], control_outputs['ans'], control_outputs_neg['ans']):\n", - " print(\"===== No Control =====\")\n", - " print(f'{s:02.2%}')\n", - " print(f\"===== + Honesty Control =====\")\n", - " print(f'{p:02.2%}')\n", - " print()\n", - " print(f\"===== - Honesty Control =====\")\n", - " print(f'{n:02.2%}')\n", - " print()\n" - ] - }, - { - "cell_type": "code", - "execution_count": 21, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "===== No Control =====\n", - "negative\n", - "===== + Honesty Control =====\n", - "negative\n", - "\n", - "===== - Honesty Control =====\n", - "I\n", - "\n", - "===== No Control =====\n", - "pos\n", - "===== + Honesty Control =====\n", - "pos\n", - "\n", - "===== - Honesty Control =====\n", - "pos\n", - "\n" - ] - } - ], + "outputs": [], "source": [ "\n", "from re import S\n", @@ -494,56 +285,28 @@ " control_outputs = rep_control_pipeline(inputs, activations=activations, batch_size=batch_size, **text_gen_kwargs)\n", " control_outputs_neg = rep_control_pipeline(inputs, activations=activations_neg, batch_size=batch_size, **text_gen_kwargs)\n", "\n", - "for i,s,p,n in zip(inputs, baseline_outputs['text_ans'], control_outputs['text_ans'], control_outputs_neg['text_ans']):\n", - " print(\"===== No Control =====\")\n", - " print(s)\n", - " print(f\"===== + Honesty Control =====\")\n", - " print(p)\n", + "names = ['No Control', '+ Honesty Control', '- Honesty Control']\n", + "for i in range(len(baseline_outputs['ans'])):\n", + " for j, r in enumerate([baseline_outputs, control_outputs, control_outputs_neg]): \n", + " choices = r['answer_choices'][i]\n", + " label = r['label_true'][i]\n", + " ans = r['ans'][i]\n", + " choice_true = choices[label]\n", + " if label==0:\n", + " ans *= -1 \n", + " print(f\"==== {names[j]} {r['example_i'][i]}====\")\n", + " print(f\"Score: {ans:02.2%} of true ans `{choice_true}`\")\n", + " # print(f\"Text ans: {r['text_ans'][i]}\") \n", " print()\n", - " print(f\"===== - Honesty Control =====\")\n", - " print(n)\n", - " print()\n" + "\n" ] }, { "cell_type": "code", - "execution_count": 22, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "===== No Control =====\n", - "15.20%\n", - "===== + Honesty Control =====\n", - "19.81%\n", - "\n", - "===== - Honesty Control =====\n", - "29.42%\n", - "\n", - "===== No Control =====\n", - "61.88%\n", - "===== + Honesty Control =====\n", - "91.96%\n", - "\n", - "===== - Honesty Control =====\n", - "84.79%\n", - "\n" - ] - } - ], - "source": [ - "for i,s,p,n in zip(inputs, baseline_outputs['ans'], control_outputs['ans'], control_outputs_neg['ans']):\n", - " print(\"===== No Control =====\")\n", - " print(f'{s:02.2%}')\n", - " print(f\"===== + Honesty Control =====\")\n", - " print(f'{p:02.2%}')\n", - " print()\n", - " print(f\"===== - Honesty Control =====\")\n", - " print(f'{n:02.2%}')\n", - " print()\n" - ] + "outputs": [], + "source": [] }, { "cell_type": "markdown", @@ -561,20 +324,9 @@ }, { "cell_type": "code", - "execution_count": 13, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "data": { - "text/plain": [ - "" - ] - }, - "execution_count": 13, - "metadata": {}, - "output_type": "execute_result" - } - ], + "outputs": [], "source": [ "\n", "\n", @@ -589,34 +341,11 @@ }, { "cell_type": "code", - "execution_count": 14, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "===== No Control =====\n", - "negative\n", - "===== + Honesty Control =====\n", - "negative\n", - "\n", - "===== - Honesty Control =====\n", - "I\n", - "\n", - "===== No Control =====\n", - "pos\n", - "===== + Honesty Control =====\n", - "pos\n", - "\n", - "===== - Honesty Control =====\n", - "pos\n", - "\n" - ] - } - ], + "outputs": [], "source": [ - "inputs = dataset_train[:2]\n", + "\n", "coeff=8.0\n", "max_new_tokens=3\n", "text_gen_kwargs = dict(do_sample=False, max_new_tokens=max_new_tokens, use_cache=False, \n", @@ -637,71 +366,21 @@ " control_outputs_neg = rep_control_pipeline2(inputs, activations=activations_neg, batch_size=batch_size, **text_gen_kwargs)\n", "\n", "\n", - "for i,s,p,n in zip(inputs, baseline_outputs['text_ans'], control_outputs['text_ans'], control_outputs_neg['text_ans']):\n", - " print(\"===== No Control =====\")\n", - " print(s)\n", - " print(f\"===== + Honesty Control =====\")\n", - " print(p)\n", - " print()\n", - " print(f\"===== - Honesty Control =====\")\n", - " print(n)\n", + "names = ['No Control', '+ Honesty Control', '- Honesty Control']\n", + "for i in range(len(baseline_outputs['ans'])):\n", + " for j, r in enumerate([baseline_outputs, control_outputs, control_outputs_neg]): \n", + " choices = r['answer_choices'][i]\n", + " label = r['label_true'][i]\n", + " ans = r['ans'][i]\n", + " choice_true = choices[label]\n", + " if label==0:\n", + " ans *= -1 \n", + " print(f\"==== {names[j]} ====\")\n", + " print(f\"Score: {ans:02.2%} of true ans `{choice_true}`\")\n", + " # print(f\"Text ans: {r['text_ans'][i]}\") \n", " print()\n" ] }, - { - "cell_type": "code", - "execution_count": 15, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "===== No Control =====\n", - "15.20%\n", - "===== + Honesty Control =====\n", - "19.81%\n", - "\n", - "===== - Honesty Control =====\n", - "29.42%\n", - "\n", - "===== No Control =====\n", - "61.88%\n", - "===== + Honesty Control =====\n", - "91.96%\n", - "\n", - "===== - Honesty Control =====\n", - "84.79%\n", - "\n" - ] - } - ], - "source": [ - "for i,s,p,n in zip(inputs, baseline_outputs['ans'], control_outputs['ans'], control_outputs_neg['ans']):\n", - " print(\"===== No Control =====\")\n", - " print(f'{s:02.2%}')\n", - " print(f\"===== + Honesty Control =====\")\n", - " print(f'{p:02.2%}')\n", - " print()\n", - " print(f\"===== - Honesty Control =====\")\n", - " print(f'{n:02.2%}')\n", - " print()\n" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [] - }, { "cell_type": "code", "execution_count": null, diff --git a/src/repe/rep_control_pipeline.py b/src/repe/rep_control_pipeline.py index 222ddd3..27d3df7 100644 --- a/src/repe/rep_control_pipeline.py +++ b/src/repe/rep_control_pipeline.py @@ -37,10 +37,10 @@ class RepControlPipeline(RepControlPipeline2): super().__init__(model=model, tokenizer=tokenizer, max_length=max_length, **kwargs) - def __call__(self, text_inputs, activations=None, **kwargs): + def __call__(self, text_inputs, activations=None, token_pos=None, masks=None, normalize=False, **kwargs): if activations is not None: self.wrapped_model.reset() - self.wrapped_model.set_controller(self.layers, activations, self.block_name) + self.wrapped_model.set_controller(self.layers, activations, self.block_name, token_pos=token_pos, masks=masks, normalize=normalize) outputs = super().__call__(text_inputs, **kwargs) self.wrapped_model.reset() diff --git a/src/repe/rep_control_reading_vec.py b/src/repe/rep_control_reading_vec.py index 7d07cb8..70e2185 100644 --- a/src/repe/rep_control_reading_vec.py +++ b/src/repe/rep_control_reading_vec.py @@ -33,11 +33,12 @@ class WrappedBlock(torch.nn.Module): # we should ignore the padding tokens when doing the activation addition # mask has ones for non padding tokens and zeros at padding tokens. # only tested this on left padding - elif "position_ids" in kwargs: + elif ("position_ids" in kwargs) and not self.token_pos: pos = kwargs["position_ids"] + pos = pos.repeat(modified.shape[0], 1, 1) zero_indices = (pos == 0).cumsum(1).argmax(1, keepdim=True) col_indices = torch.arange(pos.size(1), device=pos.device).unsqueeze(0) - target_shape = pos.shape + target_shape = modified.shape mask = (col_indices >= zero_indices).float().reshape(target_shape[0], target_shape[1], 1) mask = mask.to(modified.dtype) else: @@ -51,6 +52,7 @@ class WrappedBlock(torch.nn.Module): self.controller = self.controller.to(modified.device) if type(mask) == torch.Tensor: mask = mask.to(modified.device) + if isinstance(self.token_pos, int): modified[:, self.token_pos] = modified[:, self.token_pos] + self.controller * mask elif isinstance(self.token_pos, list) or isinstance(self.token_pos, tuple) or isinstance(self.token_pos, np.ndarray):