This commit is contained in:
deep1
2023-10-15 18:19:13 +08:00
parent 5bac922278
commit c1fff539d7
5 changed files with 342 additions and 178 deletions
+166 -68
View File
@@ -50,9 +50,9 @@
"name": "stderr",
"output_type": "stream",
"text": [
"\u001b[32m2023-10-15 16:46:07.158\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",
"\u001b[32m2023-10-15 16:46:07.159\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",
"\u001b[32m2023-10-15 16:46:07.159\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"
"\u001b[32m2023-10-15 17:26:06.435\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",
"\u001b[32m2023-10-15 17:26:06.435\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",
"\u001b[32m2023-10-15 17:26:06.436\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"
]
}
],
@@ -68,7 +68,28 @@
{
"data": {
"application/vnd.jupyter.widget-view+json": {
"model_id": "d937549193694407bccbe69577fc057e",
"model_id": "91397fa193d244de85d0b1cefbe96976",
"version_major": 2,
"version_minor": 0
},
"text/plain": [
"Generating train split: 0 examples [00:00, ? examples/s]"
]
},
"metadata": {},
"output_type": "display_data"
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Extracting 11 variants of each prompt\n"
]
},
{
"data": {
"application/vnd.jupyter.widget-view+json": {
"model_id": "c2293e5733f94538835fb1bf52c82d52",
"version_major": 2,
"version_minor": 0
},
@@ -82,7 +103,7 @@
{
"data": {
"application/vnd.jupyter.widget-view+json": {
"model_id": "11563ae0472e4215982ba483b7710715",
"model_id": "47190779321c41f3857a5b270f2d1c2c",
"version_major": 2,
"version_minor": 0
},
@@ -96,7 +117,7 @@
{
"data": {
"application/vnd.jupyter.widget-view+json": {
"model_id": "2b24eb7b7cfa42c1840a7676cbe1f942",
"model_id": "2bc743f5d3e4455086c003c3e8537cda",
"version_major": 2,
"version_minor": 0
},
@@ -110,7 +131,7 @@
{
"data": {
"application/vnd.jupyter.widget-view+json": {
"model_id": "3174bc1f2a034d5ab09d37549858ecc3",
"model_id": "7b14eb6d6e134eb49d8df9e79ec9aee5",
"version_major": 2,
"version_minor": 0
},
@@ -124,7 +145,7 @@
{
"data": {
"application/vnd.jupyter.widget-view+json": {
"model_id": "9aa208807eea47bdb460d3120b9f3953",
"model_id": "a8be2f155ad244438b9a319b51421a99",
"version_major": 2,
"version_minor": 0
},
@@ -185,13 +206,20 @@
},
{
"cell_type": "code",
"execution_count": 7,
"execution_count": null,
"metadata": {},
"outputs": [],
"source": []
},
{
"cell_type": "code",
"execution_count": 24,
"metadata": {},
"outputs": [
{
"data": {
"application/vnd.jupyter.widget-view+json": {
"model_id": "f092c8b927474979a94c3357505e4c35",
"model_id": "97d9267e2ff1497488f7951017fe54b0",
"version_major": 2,
"version_minor": 0
},
@@ -205,7 +233,7 @@
{
"data": {
"application/vnd.jupyter.widget-view+json": {
"model_id": "edfbe70a3653465da59ca5e945bce2a9",
"model_id": "9973a2a52a7a43a38c84ea2b1e4ec2fa",
"version_major": 2,
"version_minor": 0
},
@@ -217,17 +245,36 @@
"output_type": "display_data"
},
{
"data": {
"text/plain": [
"Dataset({\n",
" features: ['scores0', 'ds_index', 'scores', 'head_activation', 'ds_string', 'example_i', 'answer', 'question', 'answer_choices', 'template_name', 'label_true', 'label_instructed', 'instructed_to_lie', 'sys_instr_name', 'truncated', 'prompt_truncated'],\n",
" num_rows: 10\n",
"})"
]
},
"execution_count": 7,
"metadata": {},
"output_type": "execute_result"
"ename": "DatasetGenerationError",
"evalue": "An error occurred while generating the dataset",
"output_type": "error",
"traceback": [
"\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
"\u001b[0;31mArrowTypeError\u001b[0m Traceback (most recent call last)",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1703\u001b[0m, in \u001b[0;36mGeneratorBasedBuilder._prepare_split_single\u001b[0;34m(self, gen_kwargs, fpath, file_format, max_shard_size, split_info, check_duplicate_keys, job_id)\u001b[0m\n\u001b[1;32m 1702\u001b[0m num_shards \u001b[39m=\u001b[39m shard_id \u001b[39m+\u001b[39m \u001b[39m1\u001b[39m\n\u001b[0;32m-> 1703\u001b[0m num_examples, num_bytes \u001b[39m=\u001b[39m writer\u001b[39m.\u001b[39;49mfinalize()\n\u001b[1;32m 1704\u001b[0m writer\u001b[39m.\u001b[39mclose()\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_writer.py:586\u001b[0m, in \u001b[0;36mArrowWriter.finalize\u001b[0;34m(self, close_stream)\u001b[0m\n\u001b[1;32m 585\u001b[0m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39mhkey_record \u001b[39m=\u001b[39m []\n\u001b[0;32m--> 586\u001b[0m \u001b[39mself\u001b[39;49m\u001b[39m.\u001b[39;49mwrite_examples_on_file()\n\u001b[1;32m 587\u001b[0m \u001b[39m# If schema is known, infer features even if no examples were written\u001b[39;00m\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_writer.py:448\u001b[0m, in \u001b[0;36mArrowWriter.write_examples_on_file\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m 444\u001b[0m batch_examples[col] \u001b[39m=\u001b[39m [\n\u001b[1;32m 445\u001b[0m row[\u001b[39m0\u001b[39m][col]\u001b[39m.\u001b[39mto_pylist()[\u001b[39m0\u001b[39m] \u001b[39mif\u001b[39;00m \u001b[39misinstance\u001b[39m(row[\u001b[39m0\u001b[39m][col], (pa\u001b[39m.\u001b[39mArray, pa\u001b[39m.\u001b[39mChunkedArray)) \u001b[39melse\u001b[39;00m row[\u001b[39m0\u001b[39m][col]\n\u001b[1;32m 446\u001b[0m \u001b[39mfor\u001b[39;00m row \u001b[39min\u001b[39;00m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39mcurrent_examples\n\u001b[1;32m 447\u001b[0m ]\n\u001b[0;32m--> 448\u001b[0m \u001b[39mself\u001b[39;49m\u001b[39m.\u001b[39;49mwrite_batch(batch_examples\u001b[39m=\u001b[39;49mbatch_examples)\n\u001b[1;32m 449\u001b[0m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39mcurrent_examples \u001b[39m=\u001b[39m []\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_writer.py:555\u001b[0m, in \u001b[0;36mArrowWriter.write_batch\u001b[0;34m(self, batch_examples, writer_batch_size)\u001b[0m\n\u001b[1;32m 554\u001b[0m typed_sequence \u001b[39m=\u001b[39m OptimizedTypedSequence(col_values, \u001b[39mtype\u001b[39m\u001b[39m=\u001b[39mcol_type, try_type\u001b[39m=\u001b[39mcol_try_type, col\u001b[39m=\u001b[39mcol)\n\u001b[0;32m--> 555\u001b[0m arrays\u001b[39m.\u001b[39mappend(pa\u001b[39m.\u001b[39;49marray(typed_sequence))\n\u001b[1;32m 556\u001b[0m inferred_features[col] \u001b[39m=\u001b[39m typed_sequence\u001b[39m.\u001b[39mget_inferred_type()\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/array.pxi:243\u001b[0m, in \u001b[0;36mpyarrow.lib.array\u001b[0;34m()\u001b[0m\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/array.pxi:110\u001b[0m, in \u001b[0;36mpyarrow.lib._handle_arrow_array_protocol\u001b[0;34m()\u001b[0m\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_writer.py:189\u001b[0m, in \u001b[0;36mTypedSequence.__arrow_array__\u001b[0;34m(self, type)\u001b[0m\n\u001b[1;32m 188\u001b[0m trying_cast_to_python_objects \u001b[39m=\u001b[39m \u001b[39mTrue\u001b[39;00m\n\u001b[0;32m--> 189\u001b[0m out \u001b[39m=\u001b[39m pa\u001b[39m.\u001b[39;49marray(cast_to_python_objects(data, only_1d_for_numpy\u001b[39m=\u001b[39;49m\u001b[39mTrue\u001b[39;49;00m))\n\u001b[1;32m 190\u001b[0m \u001b[39m# use smaller integer precisions if possible\u001b[39;00m\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/array.pxi:327\u001b[0m, in \u001b[0;36mpyarrow.lib.array\u001b[0;34m()\u001b[0m\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/array.pxi:39\u001b[0m, in \u001b[0;36mpyarrow.lib._sequence_to_array\u001b[0;34m()\u001b[0m\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/error.pxi:144\u001b[0m, in \u001b[0;36mpyarrow.lib.pyarrow_internal_check_status\u001b[0;34m()\u001b[0m\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/error.pxi:123\u001b[0m, in \u001b[0;36mpyarrow.lib.check_status\u001b[0;34m()\u001b[0m\n",
"\u001b[0;31mArrowTypeError\u001b[0m: Expected bytes, got a 'list' object",
"\nThe above exception was the direct cause of the following exception:\n",
"\u001b[0;31mDatasetGenerationError\u001b[0m Traceback (most recent call last)",
"\u001b[1;32m/home/ubuntu/Documents/mjc/elk/discovering_latent_knowledge2/notebooks/012b_scratch_dataset.ipynb Cell 9\u001b[0m line \u001b[0;36m1\n\u001b[0;32m----> <a href='vscode-notebook-cell://ssh-remote%2Bdeep1-local/home/ubuntu/Documents/mjc/elk/discovering_latent_knowledge2/notebooks/012b_scratch_dataset.ipynb#X13sdnNjb2RlLXJlbW90ZQ%3D%3D?line=0'>1</a>\u001b[0m ds1 \u001b[39m=\u001b[39m Dataset\u001b[39m.\u001b[39;49mfrom_generator(\n\u001b[1;32m <a href='vscode-notebook-cell://ssh-remote%2Bdeep1-local/home/ubuntu/Documents/mjc/elk/discovering_latent_knowledge2/notebooks/012b_scratch_dataset.ipynb#X13sdnNjb2RlLXJlbW90ZQ%3D%3D?line=1'>2</a>\u001b[0m generator\u001b[39m=\u001b[39;49mbatch_hidden_states,\n\u001b[1;32m <a href='vscode-notebook-cell://ssh-remote%2Bdeep1-local/home/ubuntu/Documents/mjc/elk/discovering_latent_knowledge2/notebooks/012b_scratch_dataset.ipynb#X13sdnNjb2RlLXJlbW90ZQ%3D%3D?line=2'>3</a>\u001b[0m info\u001b[39m=\u001b[39;49mDatasetInfo(\n\u001b[1;32m <a href='vscode-notebook-cell://ssh-remote%2Bdeep1-local/home/ubuntu/Documents/mjc/elk/discovering_latent_knowledge2/notebooks/012b_scratch_dataset.ipynb#X13sdnNjb2RlLXJlbW90ZQ%3D%3D?line=3'>4</a>\u001b[0m description\u001b[39m=\u001b[39;49mjson\u001b[39m.\u001b[39;49mdumps(info_kwargs, indent\u001b[39m=\u001b[39;49m\u001b[39m2\u001b[39;49m),\n\u001b[1;32m <a href='vscode-notebook-cell://ssh-remote%2Bdeep1-local/home/ubuntu/Documents/mjc/elk/discovering_latent_knowledge2/notebooks/012b_scratch_dataset.ipynb#X13sdnNjb2RlLXJlbW90ZQ%3D%3D?line=4'>5</a>\u001b[0m config_name\u001b[39m=\u001b[39;49mf,\n\u001b[1;32m <a href='vscode-notebook-cell://ssh-remote%2Bdeep1-local/home/ubuntu/Documents/mjc/elk/discovering_latent_knowledge2/notebooks/012b_scratch_dataset.ipynb#X13sdnNjb2RlLXJlbW90ZQ%3D%3D?line=5'>6</a>\u001b[0m ),\n\u001b[1;32m <a href='vscode-notebook-cell://ssh-remote%2Bdeep1-local/home/ubuntu/Documents/mjc/elk/discovering_latent_knowledge2/notebooks/012b_scratch_dataset.ipynb#X13sdnNjb2RlLXJlbW90ZQ%3D%3D?line=6'>7</a>\u001b[0m gen_kwargs\u001b[39m=\u001b[39;49mgen_kwargs,\n\u001b[1;32m <a href='vscode-notebook-cell://ssh-remote%2Bdeep1-local/home/ubuntu/Documents/mjc/elk/discovering_latent_knowledge2/notebooks/012b_scratch_dataset.ipynb#X13sdnNjb2RlLXJlbW90ZQ%3D%3D?line=7'>8</a>\u001b[0m num_proc\u001b[39m=\u001b[39;49m\u001b[39m1\u001b[39;49m,\n\u001b[1;32m <a href='vscode-notebook-cell://ssh-remote%2Bdeep1-local/home/ubuntu/Documents/mjc/elk/discovering_latent_knowledge2/notebooks/012b_scratch_dataset.ipynb#X13sdnNjb2RlLXJlbW90ZQ%3D%3D?line=8'>9</a>\u001b[0m )\n\u001b[1;32m <a href='vscode-notebook-cell://ssh-remote%2Bdeep1-local/home/ubuntu/Documents/mjc/elk/discovering_latent_knowledge2/notebooks/012b_scratch_dataset.ipynb#X13sdnNjb2RlLXJlbW90ZQ%3D%3D?line=9'>10</a>\u001b[0m ds1\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_dataset.py:1072\u001b[0m, in \u001b[0;36mDataset.from_generator\u001b[0;34m(generator, features, cache_dir, keep_in_memory, gen_kwargs, num_proc, **kwargs)\u001b[0m\n\u001b[1;32m 1016\u001b[0m \u001b[39m\u001b[39m\u001b[39m\"\"\"Create a Dataset from a generator.\u001b[39;00m\n\u001b[1;32m 1017\u001b[0m \n\u001b[1;32m 1018\u001b[0m \u001b[39mArgs:\u001b[39;00m\n\u001b[0;32m (...)\u001b[0m\n\u001b[1;32m 1060\u001b[0m \u001b[39m```\u001b[39;00m\n\u001b[1;32m 1061\u001b[0m \u001b[39m\"\"\"\u001b[39;00m\n\u001b[1;32m 1062\u001b[0m \u001b[39mfrom\u001b[39;00m \u001b[39m.\u001b[39;00m\u001b[39mio\u001b[39;00m\u001b[39m.\u001b[39;00m\u001b[39mgenerator\u001b[39;00m \u001b[39mimport\u001b[39;00m GeneratorDatasetInputStream\n\u001b[1;32m 1064\u001b[0m \u001b[39mreturn\u001b[39;00m GeneratorDatasetInputStream(\n\u001b[1;32m 1065\u001b[0m generator\u001b[39m=\u001b[39;49mgenerator,\n\u001b[1;32m 1066\u001b[0m features\u001b[39m=\u001b[39;49mfeatures,\n\u001b[1;32m 1067\u001b[0m cache_dir\u001b[39m=\u001b[39;49mcache_dir,\n\u001b[1;32m 1068\u001b[0m keep_in_memory\u001b[39m=\u001b[39;49mkeep_in_memory,\n\u001b[1;32m 1069\u001b[0m gen_kwargs\u001b[39m=\u001b[39;49mgen_kwargs,\n\u001b[1;32m 1070\u001b[0m num_proc\u001b[39m=\u001b[39;49mnum_proc,\n\u001b[1;32m 1071\u001b[0m \u001b[39m*\u001b[39;49m\u001b[39m*\u001b[39;49mkwargs,\n\u001b[0;32m-> 1072\u001b[0m )\u001b[39m.\u001b[39;49mread()\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/io/generator.py:47\u001b[0m, in \u001b[0;36mGeneratorDatasetInputStream.read\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m 44\u001b[0m verification_mode \u001b[39m=\u001b[39m \u001b[39mNone\u001b[39;00m\n\u001b[1;32m 45\u001b[0m base_path \u001b[39m=\u001b[39m \u001b[39mNone\u001b[39;00m\n\u001b[0;32m---> 47\u001b[0m \u001b[39mself\u001b[39;49m\u001b[39m.\u001b[39;49mbuilder\u001b[39m.\u001b[39;49mdownload_and_prepare(\n\u001b[1;32m 48\u001b[0m download_config\u001b[39m=\u001b[39;49mdownload_config,\n\u001b[1;32m 49\u001b[0m download_mode\u001b[39m=\u001b[39;49mdownload_mode,\n\u001b[1;32m 50\u001b[0m verification_mode\u001b[39m=\u001b[39;49mverification_mode,\n\u001b[1;32m 51\u001b[0m \u001b[39m# try_from_hf_gcs=try_from_hf_gcs,\u001b[39;49;00m\n\u001b[1;32m 52\u001b[0m base_path\u001b[39m=\u001b[39;49mbase_path,\n\u001b[1;32m 53\u001b[0m num_proc\u001b[39m=\u001b[39;49m\u001b[39mself\u001b[39;49m\u001b[39m.\u001b[39;49mnum_proc,\n\u001b[1;32m 54\u001b[0m )\n\u001b[1;32m 55\u001b[0m dataset \u001b[39m=\u001b[39m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39mbuilder\u001b[39m.\u001b[39mas_dataset(\n\u001b[1;32m 56\u001b[0m split\u001b[39m=\u001b[39m\u001b[39m\"\u001b[39m\u001b[39mtrain\u001b[39m\u001b[39m\"\u001b[39m, verification_mode\u001b[39m=\u001b[39mverification_mode, in_memory\u001b[39m=\u001b[39m\u001b[39mself\u001b[39m\u001b[39m.\u001b[39mkeep_in_memory\n\u001b[1;32m 57\u001b[0m )\n\u001b[1;32m 58\u001b[0m \u001b[39mreturn\u001b[39;00m dataset\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:954\u001b[0m, in \u001b[0;36mDatasetBuilder.download_and_prepare\u001b[0;34m(self, output_dir, download_config, download_mode, verification_mode, ignore_verifications, try_from_hf_gcs, dl_manager, base_path, use_auth_token, file_format, max_shard_size, num_proc, storage_options, **download_and_prepare_kwargs)\u001b[0m\n\u001b[1;32m 952\u001b[0m \u001b[39mif\u001b[39;00m num_proc \u001b[39mis\u001b[39;00m \u001b[39mnot\u001b[39;00m \u001b[39mNone\u001b[39;00m:\n\u001b[1;32m 953\u001b[0m prepare_split_kwargs[\u001b[39m\"\u001b[39m\u001b[39mnum_proc\u001b[39m\u001b[39m\"\u001b[39m] \u001b[39m=\u001b[39m num_proc\n\u001b[0;32m--> 954\u001b[0m \u001b[39mself\u001b[39;49m\u001b[39m.\u001b[39;49m_download_and_prepare(\n\u001b[1;32m 955\u001b[0m dl_manager\u001b[39m=\u001b[39;49mdl_manager,\n\u001b[1;32m 956\u001b[0m verification_mode\u001b[39m=\u001b[39;49mverification_mode,\n\u001b[1;32m 957\u001b[0m \u001b[39m*\u001b[39;49m\u001b[39m*\u001b[39;49mprepare_split_kwargs,\n\u001b[1;32m 958\u001b[0m \u001b[39m*\u001b[39;49m\u001b[39m*\u001b[39;49mdownload_and_prepare_kwargs,\n\u001b[1;32m 959\u001b[0m )\n\u001b[1;32m 960\u001b[0m \u001b[39m# Sync info\u001b[39;00m\n\u001b[1;32m 961\u001b[0m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39minfo\u001b[39m.\u001b[39mdataset_size \u001b[39m=\u001b[39m \u001b[39msum\u001b[39m(split\u001b[39m.\u001b[39mnum_bytes \u001b[39mfor\u001b[39;00m split \u001b[39min\u001b[39;00m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39minfo\u001b[39m.\u001b[39msplits\u001b[39m.\u001b[39mvalues())\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1717\u001b[0m, in \u001b[0;36mGeneratorBasedBuilder._download_and_prepare\u001b[0;34m(self, dl_manager, verification_mode, **prepare_splits_kwargs)\u001b[0m\n\u001b[1;32m 1716\u001b[0m \u001b[39mdef\u001b[39;00m \u001b[39m_download_and_prepare\u001b[39m(\u001b[39mself\u001b[39m, dl_manager, verification_mode, \u001b[39m*\u001b[39m\u001b[39m*\u001b[39mprepare_splits_kwargs):\n\u001b[0;32m-> 1717\u001b[0m \u001b[39msuper\u001b[39;49m()\u001b[39m.\u001b[39;49m_download_and_prepare(\n\u001b[1;32m 1718\u001b[0m dl_manager,\n\u001b[1;32m 1719\u001b[0m verification_mode,\n\u001b[1;32m 1720\u001b[0m check_duplicate_keys\u001b[39m=\u001b[39;49mverification_mode \u001b[39m==\u001b[39;49m VerificationMode\u001b[39m.\u001b[39;49mBASIC_CHECKS\n\u001b[1;32m 1721\u001b[0m \u001b[39mor\u001b[39;49;00m verification_mode \u001b[39m==\u001b[39;49m VerificationMode\u001b[39m.\u001b[39;49mALL_CHECKS,\n\u001b[1;32m 1722\u001b[0m \u001b[39m*\u001b[39;49m\u001b[39m*\u001b[39;49mprepare_splits_kwargs,\n\u001b[1;32m 1723\u001b[0m )\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1049\u001b[0m, in \u001b[0;36mDatasetBuilder._download_and_prepare\u001b[0;34m(self, dl_manager, verification_mode, **prepare_split_kwargs)\u001b[0m\n\u001b[1;32m 1045\u001b[0m split_dict\u001b[39m.\u001b[39madd(split_generator\u001b[39m.\u001b[39msplit_info)\n\u001b[1;32m 1047\u001b[0m \u001b[39mtry\u001b[39;00m:\n\u001b[1;32m 1048\u001b[0m \u001b[39m# Prepare split will record examples associated to the split\u001b[39;00m\n\u001b[0;32m-> 1049\u001b[0m \u001b[39mself\u001b[39;49m\u001b[39m.\u001b[39;49m_prepare_split(split_generator, \u001b[39m*\u001b[39;49m\u001b[39m*\u001b[39;49mprepare_split_kwargs)\n\u001b[1;32m 1050\u001b[0m \u001b[39mexcept\u001b[39;00m \u001b[39mOSError\u001b[39;00m \u001b[39mas\u001b[39;00m e:\n\u001b[1;32m 1051\u001b[0m \u001b[39mraise\u001b[39;00m \u001b[39mOSError\u001b[39;00m(\n\u001b[1;32m 1052\u001b[0m \u001b[39m\"\u001b[39m\u001b[39mCannot find data file. \u001b[39m\u001b[39m\"\u001b[39m\n\u001b[1;32m 1053\u001b[0m \u001b[39m+\u001b[39m (\u001b[39mself\u001b[39m\u001b[39m.\u001b[39mmanual_download_instructions \u001b[39mor\u001b[39;00m \u001b[39m\"\u001b[39m\u001b[39m\"\u001b[39m)\n\u001b[1;32m 1054\u001b[0m \u001b[39m+\u001b[39m \u001b[39m\"\u001b[39m\u001b[39m\\n\u001b[39;00m\u001b[39mOriginal error:\u001b[39m\u001b[39m\\n\u001b[39;00m\u001b[39m\"\u001b[39m\n\u001b[1;32m 1055\u001b[0m \u001b[39m+\u001b[39m \u001b[39mstr\u001b[39m(e)\n\u001b[1;32m 1056\u001b[0m ) \u001b[39mfrom\u001b[39;00m \u001b[39mNone\u001b[39;00m\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1555\u001b[0m, in \u001b[0;36mGeneratorBasedBuilder._prepare_split\u001b[0;34m(self, split_generator, check_duplicate_keys, file_format, num_proc, max_shard_size)\u001b[0m\n\u001b[1;32m 1553\u001b[0m job_id \u001b[39m=\u001b[39m \u001b[39m0\u001b[39m\n\u001b[1;32m 1554\u001b[0m \u001b[39mwith\u001b[39;00m pbar:\n\u001b[0;32m-> 1555\u001b[0m \u001b[39mfor\u001b[39;00m job_id, done, content \u001b[39min\u001b[39;00m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39m_prepare_split_single(\n\u001b[1;32m 1556\u001b[0m gen_kwargs\u001b[39m=\u001b[39mgen_kwargs, job_id\u001b[39m=\u001b[39mjob_id, \u001b[39m*\u001b[39m\u001b[39m*\u001b[39m_prepare_split_args\n\u001b[1;32m 1557\u001b[0m ):\n\u001b[1;32m 1558\u001b[0m \u001b[39mif\u001b[39;00m done:\n\u001b[1;32m 1559\u001b[0m result \u001b[39m=\u001b[39m content\n",
"File \u001b[0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1712\u001b[0m, in \u001b[0;36mGeneratorBasedBuilder._prepare_split_single\u001b[0;34m(self, gen_kwargs, fpath, file_format, max_shard_size, split_info, check_duplicate_keys, job_id)\u001b[0m\n\u001b[1;32m 1710\u001b[0m \u001b[39mif\u001b[39;00m \u001b[39misinstance\u001b[39m(e, SchemaInferenceError) \u001b[39mand\u001b[39;00m e\u001b[39m.\u001b[39m__context__ \u001b[39mis\u001b[39;00m \u001b[39mnot\u001b[39;00m \u001b[39mNone\u001b[39;00m:\n\u001b[1;32m 1711\u001b[0m e \u001b[39m=\u001b[39m e\u001b[39m.\u001b[39m__context__\n\u001b[0;32m-> 1712\u001b[0m \u001b[39mraise\u001b[39;00m DatasetGenerationError(\u001b[39m\"\u001b[39m\u001b[39mAn error occurred while generating the dataset\u001b[39m\u001b[39m\"\u001b[39m) \u001b[39mfrom\u001b[39;00m \u001b[39me\u001b[39;00m\n\u001b[1;32m 1714\u001b[0m \u001b[39myield\u001b[39;00m job_id, \u001b[39mTrue\u001b[39;00m, (total_num_examples, total_num_bytes, writer\u001b[39m.\u001b[39m_features, num_shards, shard_lengths)\n",
"\u001b[0;31mDatasetGenerationError\u001b[0m: An error occurred while generating the dataset"
]
}
],
"source": [
@@ -243,6 +290,36 @@
"ds1\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"['model.layers.4.self_attn',\n",
" 'model.layers.8.self_attn',\n",
" 'model.layers.4.self_attn',\n",
" 'model.layers.8.self_attn',\n",
" 'model.layers.4.self_attn',\n",
" 'model.layers.8.self_attn',\n",
" 'model.layers.4.self_attn',\n",
" 'model.layers.8.self_attn',\n",
" 'model.layers.4.self_attn',\n",
" 'model.layers.8.self_attn']"
]
},
"execution_count": 23,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"ds1.info.description\n",
"ds1['layer_names']\n"
]
},
{
"cell_type": "markdown",
"metadata": {},
@@ -305,7 +382,7 @@
{
"data": {
"application/vnd.jupyter.widget-view+json": {
"model_id": "6f49fdcc67b4496fbdd136c8b10dd024",
"model_id": "439e450b8f224791a879a3adc7696d4f",
"version_major": 2,
"version_minor": 0
},
@@ -365,6 +442,9 @@
"from einops import rearrange, reduce, repeat, asnumpy, parse_shape\n",
"from src.datasets.intervene import InterventionDict\n",
"from typing import Tuple\n",
"from functools import partial\n",
"from baukit.nethook import Trace, TraceDict, recursive_copy\n",
"from src.datasets.intervene import intervention_meta_fn, get_interventions_dict\n",
"\n",
"activations = np.array(ds1['head_activation']).squeeze(-1)\n",
"labels = np.array(ds1[\"label_true\"]).astype(int)==1\n",
@@ -376,39 +456,57 @@
},
{
"cell_type": "code",
"execution_count": 14,
"execution_count": 12,
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"tensor([[[-3.4062, 12.5859, 0.4016, ..., -2.1797, -1.2871, -2.1543],\n",
" [-3.4082, 12.5625, 0.3992, ..., -2.1777, -1.2900, -2.1543],\n",
" [-3.4062, 12.5703, 0.3999, ..., -2.1777, -1.2910, -2.1543],\n",
" ...,\n",
" [-2.4629, 0.2039, 12.3828, ..., -2.2266, 0.8071, 0.2996],\n",
" [-6.0508, -6.7305, 10.3047, ..., -4.7031, -2.3340, -2.0586],\n",
" [-7.2188, -8.3750, 9.3047, ..., -5.4141, -2.8828, -3.3223]],\n",
"\n",
" [[-3.4004, 12.8203, 0.4661, ..., -2.1914, -1.1826, -2.1445],\n",
" [-3.4023, 12.7891, 0.4626, ..., -2.1914, -1.1885, -2.1465],\n",
" [-3.4004, 12.7969, 0.4636, ..., -2.1914, -1.1865, -2.1445],\n",
" ...,\n",
" [-1.5488, 0.2537, 10.4453, ..., -1.8154, 0.9443, 0.6328],\n",
" [-6.2344, -7.9570, 7.3828, ..., -4.1797, -3.3203, -2.4883],\n",
" [-5.7266, -7.7461, 10.3828, ..., -3.4277, -2.1797, -2.3887]]],\n",
" device='cuda:0')"
"'{\\n \"extract_cfg\": {\\n \"datasets\": [\\n \"amazon_polarity\",\\n \"super_glue:boolq\",\\n \"glue:qnli\",\\n \"imdb\"\\n ],\\n \"model\": \"TheBloke/WizardCoder-Python-13B-V1.0-GPTQ\",\\n \"data_dirs\": [],\\n \"max_examples\": [\\n 10,\\n 10\\n ],\\n \"num_shots\": 1,\\n \"num_variants\": -1,\\n \"layers\": [],\\n \"seed\": 42,\\n \"token_loc\": \"last\",\\n \"template_path\": null,\\n \"max_length\": 999\\n },\\n \"ds_name\": \"amazon_polarity\",\\n \"split_type\": \"train\",\\n \"f\": null,\\n \"date\": \"2023-10-15T17:26:26.289433\"\\n}'"
]
},
"execution_count": 14,
"execution_count": 12,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"from functools import partial\n",
"from baukit.nethook import Trace, TraceDict, recursive_copy\n",
"from src.datasets.intervene import intervention_meta_fn, get_interventions_dict\n",
"ds1.info.description\n"
]
},
{
"cell_type": "code",
"execution_count": 13,
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"tensor([[[-3.4297, 12.6875, 0.4692, ..., -2.2012, -1.2109, -2.1602],\n",
" [-3.4297, 12.6875, 0.4695, ..., -2.2012, -1.2129, -2.1621],\n",
" [-3.4336, 12.6875, 0.4702, ..., -2.2051, -1.2129, -2.1621],\n",
" ...,\n",
" [-2.4434, 1.2295, 11.8359, ..., -2.0957, 0.1870, 0.2942],\n",
" [-4.2383, -4.3594, 12.7656, ..., -2.8594, -1.1748, -0.6260],\n",
" [-6.4336, -6.3516, 11.0547, ..., -5.6250, -2.6367, -2.8652]],\n",
"\n",
" [[-3.4062, 12.5859, 0.4016, ..., -2.1797, -1.2871, -2.1543],\n",
" [-3.4082, 12.5625, 0.3992, ..., -2.1777, -1.2900, -2.1543],\n",
" [-3.4062, 12.5703, 0.3999, ..., -2.1777, -1.2910, -2.1543],\n",
" ...,\n",
" [-2.4629, 0.2039, 12.3828, ..., -2.2266, 0.8071, 0.2996],\n",
" [-6.0508, -6.7305, 10.3047, ..., -4.7031, -2.3340, -2.0586],\n",
" [-7.2188, -8.3750, 9.3047, ..., -5.4141, -2.8828, -3.3223]]],\n",
" device='cuda:0')"
]
},
"execution_count": 13,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"\n",
"\n",
"interventions = get_interventions_dict(activations, labels, layer_names, num_heads)\n",
"intervention_fn = partial(intervention_meta_fn, interventions=interventions, num_heads=num_heads)\n",
@@ -424,31 +522,31 @@
},
{
"cell_type": "code",
"execution_count": 15,
"execution_count": 14,
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"tensor([[[-3.4043, 12.5703, 0.4009, ..., -2.1777, -1.2881, -2.1523],\n",
"tensor([[[-3.4277, 12.6953, 0.4692, ..., -2.2012, -1.2100, -2.1602],\n",
" [-3.4297, 12.6875, 0.4695, ..., -2.2012, -1.2119, -2.1602],\n",
" [-3.4297, 12.6953, 0.4705, ..., -2.2031, -1.2119, -2.1602],\n",
" ...,\n",
" [-2.4434, 1.2295, 11.8359, ..., -2.0957, 0.1870, 0.2942],\n",
" [-4.2383, -4.3594, 12.7656, ..., -2.8594, -1.1748, -0.6260],\n",
" [-5.4531, -4.9375, 10.1641, ..., -4.8320, -2.1133, -2.2773]],\n",
"\n",
" [[-3.4043, 12.5703, 0.4009, ..., -2.1777, -1.2881, -2.1523],\n",
" [-3.4043, 12.5781, 0.4006, ..., -2.1758, -1.2881, -2.1523],\n",
" [-3.4043, 12.5703, 0.4006, ..., -2.1777, -1.2900, -2.1543],\n",
" ...,\n",
" [-2.4629, 0.2039, 12.3828, ..., -2.2266, 0.8071, 0.2996],\n",
" [-6.0508, -6.7305, 10.3047, ..., -4.7031, -2.3340, -2.0586],\n",
" [-6.3906, -7.0547, 8.7266, ..., -4.8984, -2.3691, -2.8887]],\n",
"\n",
" [[-3.4023, 12.8203, 0.4658, ..., -2.1934, -1.1855, -2.1445],\n",
" [-3.4023, 12.7891, 0.4644, ..., -2.1914, -1.1875, -2.1445],\n",
" [-3.4023, 12.8125, 0.4648, ..., -2.1934, -1.1865, -2.1445],\n",
" ...,\n",
" [-1.5488, 0.2537, 10.4453, ..., -1.8154, 0.9443, 0.6328],\n",
" [-6.2344, -7.9570, 7.3828, ..., -4.1797, -3.3203, -2.4883],\n",
" [-4.9375, -6.7891, 9.5312, ..., -3.1250, -1.7793, -1.9893]]],\n",
" [-6.3906, -7.0547, 8.7266, ..., -4.8984, -2.3691, -2.8887]]],\n",
" device='cuda:0')"
]
},
"execution_count": 15,
"execution_count": 14,
"metadata": {},
"output_type": "execute_result"
}
@@ -463,31 +561,31 @@
},
{
"cell_type": "code",
"execution_count": 22,
"execution_count": 15,
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"tensor([[[-3.4062, 12.5859, 0.4014, ..., -2.1797, -1.2881, -2.1543],\n",
"tensor([[[-3.4277, 12.6953, 0.4695, ..., -2.2012, -1.2090, -2.1602],\n",
" [-3.4277, 12.6875, 0.4697, ..., -2.1992, -1.2119, -2.1602],\n",
" [-3.4316, 12.6875, 0.4705, ..., -2.2051, -1.2119, -2.1621],\n",
" ...,\n",
" [-2.4434, 1.2295, 11.8359, ..., -2.0957, 0.1870, 0.2942],\n",
" [-4.2383, -4.3594, 12.7656, ..., -2.8594, -1.1748, -0.6260],\n",
" [-4.9805, -4.4336, 8.6094, ..., -4.4102, -1.9639, -2.0254]],\n",
"\n",
" [[-3.4062, 12.5859, 0.4014, ..., -2.1797, -1.2881, -2.1543],\n",
" [-3.4082, 12.5859, 0.4011, ..., -2.1797, -1.2881, -2.1562],\n",
" [-3.4062, 12.5703, 0.3999, ..., -2.1777, -1.2910, -2.1543],\n",
" ...,\n",
" [-2.4629, 0.2039, 12.3828, ..., -2.2266, 0.8071, 0.2996],\n",
" [-6.0508, -6.7305, 10.3047, ..., -4.7031, -2.3340, -2.0586],\n",
" [-5.8555, -6.3125, 7.6172, ..., -4.5625, -2.0273, -2.6055]],\n",
"\n",
" [[-3.4023, 12.8125, 0.4661, ..., -2.1953, -1.1846, -2.1445],\n",
" [-3.4023, 12.7891, 0.4639, ..., -2.1914, -1.1885, -2.1445],\n",
" [-3.4004, 12.7969, 0.4629, ..., -2.1895, -1.1885, -2.1445],\n",
" ...,\n",
" [-1.5488, 0.2537, 10.4453, ..., -1.8154, 0.9443, 0.6328],\n",
" [-6.2344, -7.9570, 7.3828, ..., -4.1797, -3.3203, -2.4883],\n",
" [-4.5547, -6.3555, 7.9883, ..., -3.1680, -1.6172, -1.8477]]],\n",
" [-5.8555, -6.3125, 7.6172, ..., -4.5625, -2.0273, -2.6055]]],\n",
" device='cuda:0')"
]
},
"execution_count": 22,
"execution_count": 15,
"metadata": {},
"output_type": "execute_result"
}
+124 -85
View File
@@ -14,7 +14,7 @@ logger.add("make_dataset_{time}.log")
import pandas as pd
import numpy as np
from typing import Optional, List, Dict, Union
from typing import Optional, List, Dict, Union, Set
import torch
import torch.nn as nn
@@ -28,28 +28,30 @@ from pathlib import Path
import transformers
from transformers import GPTQConfig
from datasets import Dataset, DatasetInfo
from src.datasets.load import load_ds
from src.models.load import verbose_change_param, AutoConfig, AutoTokenizer, AutoModelForCausalLM, PreTrainedTokenizerBase
from tqdm.auto import tqdm
import os, re, sys, collections, functools, itertools, json
from simple_parsing import ArgumentParser
import random
from einops import rearrange, reduce, repeat, asnumpy, parse_shape
from functools import partial
from src.datasets.load import load_ds
from src.models.load import verbose_change_param, AutoConfig, AutoTokenizer, AutoModelForCausalLM, PreTrainedTokenizerBase
from src.models.load import load_model
from src.datasets.load import ds2df
from src.datasets.load import rows_item
from src.datasets.batch import batch_hidden_states
# from src.datasets.scores import choice2ids, scores2choice_probs
from src.datasets.hs import ExtractHiddenStates
from itertools import chain
import functools
from src.prompts.prompt_loading import load_prompts
from src.datasets.scores import scores2choice_probs
from src.datasets.scores import choice2id, choice2ids
from simple_parsing import ArgumentParser
from src.datasets.intervene import intervention_meta_fn, get_interventions_dict, InterventionDict
from src.extraction.config import ExtractConfig
import random
from src.config import root_folder
def qc_ds(f):
ds4 = load_ds(f)
@@ -221,8 +223,9 @@ def qc_ds(f):
def load_preproc_dataset(ds_name: str, cfg: ExtractConfig, tokenizer: PreTrainedTokenizerBase, split_type:str="train") -> Dataset:
N = cfg.max_examples[split_type!="train"]
def load_preproc_dataset(ds_name: str, cfg: ExtractConfig, tokenizer: PreTrainedTokenizerBase, split_type:str="train", N=None) -> Dataset:
if N is None:
N = cfg.max_examples[split_type!="train"]
ds_prompts = Dataset.from_generator(
load_prompts,
gen_kwargs=dict(
@@ -272,6 +275,111 @@ def load_preproc_dataset(ds_name: str, cfg: ExtractConfig, tokenizer: PreTrained
def row_choice_ids(r, tokenizer):
return choice2ids([[c] for c in r['answer_choices']], tokenizer)
def expand_choices(choices: List[str]) -> Set[str]:
"""expand out choices by adding versions that are upper, lower, whitespace, etc"""
new = []
for c in choices:
new.append(c)
new.append(c.upper())
new.append(c.capitalize())
new.append(c.lower())
return set(new)
def post_proc_hs_ds(ds1, tokenizer):
"""add labels for our probe.
Given next_token scores (logits) we take only the subset the corresponds to our negative tokens (e.g. False, no, ...) and positive tokens (e.g. Yes, yes, affirmative, ...)."""
# left_choices = list(r[0] for r in ds1['answer_choices'])+['no', 'false', 'negative', 'wrong']
# right_choices = list(r[1] for r in ds1['answer_choices'])+['yes', 'true', 'positive', 'right']
# left_choices, right_choices = expand_choices(left_choices), expand_choices(right_choices)
# assert len(set(left_choices).intersection(right_choices))==0
# expanded_choices = [left_choices, right_choices]
# expanded_choice_ids = choice2ids(expanded_choices, tokenizer)
# add_ans_exp = lambda r: scores2choice_probs(r, expanded_choice_ids, prefix="expanded_", keys=["scores0"])
# FIXME: we have the yes and no swapped
# FIXME: use k-mean closest tokens?
# this is just based on pairs for that answer...
# FIXME of course I added a dim!
add_txt_ans0 = lambda r: {'txt_ans0': tokenizer.batch_decode(torch.softmax(torch.tensor(r['scores0']), 0).argmax(0))}
# Either just use the template choices
add_ans = lambda r: scores2choice_probs(r, row_choice_ids(r, tokenizer), keys=["scores0"])
# Or all expanded choices
ds1.set_format(type='numpy')#, columns=['input_ids', 'token_type_ids', 'attention_mask', 'label'])
ds3 = (
ds1
.map(add_ans, desc='add_ans') # slow?
# .map(add_ans_exp)
.map(add_txt_ans0, desc='add_txt_ans0')
)
return ds3
def create_hs_ds(ds_name, ds_tokens, model, cfg, intervention_dicts = [None, ], f = None, split_type="train"):
info_kwargs = dict(extract_cfg=cfg.to_dict(), ds_name=ds_name, split_type=split_type, f=f, date=pd.Timestamp.now().isoformat(),)
# first we make the calibration dataset with no intervention
gen_kwargs = dict(
model=model,
tokenizer=tokenizer,
data=ds_tokens,
batch_size=BATCH_SIZE,
layer_padding=cfg.layer_padding,
layer_stride=cfg.layer_stride,
intervention_dicts=intervention_dicts,
)
if os.environ.get('TEST', False):
# it's easier to debug if we don't use multiprocessing
gen = batch_hidden_states(**gen_kwargs)
b =next(iter(gen))
ds1 = Dataset.from_generator(
generator=batch_hidden_states,
info=DatasetInfo(
description=json.dumps(info_kwargs, indent=2),
config_name=f,
),
gen_kwargs=gen_kwargs,
num_proc=1,
)
return ds1
def create_intervention(ds_name, ds_tokens, model, layer_names, N=10):
ds_tokens_calib = ds_tokens.select(range(N-1))
# TODO: do we need ds_name if we have the ds?
ds_calibration = create_hs_ds(ds_name+'_calib', ds_tokens_calib, model, cfg, intervention_dicts = None, f=f)
activations = np.array(ds_calibration['head_activation']).squeeze(-1)
labels = np.array(ds_calibration["label_true"]).astype(int)==1
num_heads = model.config.num_attention_heads
interventions = get_interventions_dict(activations, labels, layer_names, num_heads)
return interventions
def load_intervention(ds_name, cfg, model, tokenizer, model_name):
num_heads = model.config.num_attention_heads
intervention_f = root_folder / 'data' / 'interventions' / f'{model_name}.pkl'
if not intervention_f.exists():
layer_names, layer_inds = ExtractHiddenStates(model, tokenizer, layer_stride=cfg.layer_stride, layer_padding=cfg.layer_padding).get_layer_names()
ds_tokens = load_preproc_dataset(ds_name, cfg, tokenizer, N=10)
interventions = create_intervention(ds_name, ds_tokens, model, layer_names)
torch.save(interventions, intervention_f)
else:
logger.info(f'loading interventions from {intervention_f}')
interventions = torch.load(intervention_f)
intervention_fn = partial(intervention_meta_fn, interventions=interventions, num_heads=num_heads)
return interventions, intervention_fn
if __name__ == "__main__":
parser = ArgumentParser(add_help=False)
parser.add_arguments(ExtractConfig, dest="run")
@@ -301,22 +409,16 @@ if __name__ == "__main__":
BATCH_SIZE = 4 # None # None means auto # 6 gives 16Gb/25GB. where 10GB is the base model. so 6 is 6/15
# TODO: loop through all prompts in this dataset
ds_names = cfg.datasets
split_type = "train"
model, tokenizer = load_model(cfg.model)
# def create_intervention(ds_name, model):
# # UPTO: FIXME:
# intervention_f = root / 'data' / 'interventions' / f'{model_name}.pkl'
# if not intervention_f.exist():
# intervention = calibrate(model, dataset.shuffle(42).head(10))
# # load intervention
# # get com direction https://github.com/likenneth/honest_llama/blob/master/utils.py#L731
# intervention = torch.load(intervention_f)
ds_name = 'imdb'
model_name = cfg.model
intervention, intervention_fn = load_intervention(ds_name, cfg, model, tokenizer, model_name)
for ds_name in ds_names:
@@ -328,73 +430,10 @@ if __name__ == "__main__":
sanitize = lambda s:s.replace('/', '').replace('-', '_') if s is not None else s
dataset_name = f"{sanitize(cfg.model)}_{ds_name}_{split_type}_{N}"
f = f"../.ds/{dataset_name}"
gen_kwargs = dict(
model=model,
tokenizer=tokenizer,
data=ds_tokens,
batch_size=BATCH_SIZE,
layer_padding=cfg.layer_padding,
layer_stride=cfg.layer_stride,
)
info_kwargs = dict(extract_cfg=cfg.to_dict(), ds_name=ds_name, split_type=split_type, f=f, date=pd.Timestamp.now().isoformat(),)
if os.environ.get('TEST', False):
gen = batch_hidden_states(**gen_kwargs)
b =next(iter(gen))
# [DatasetInfo](https://github.com/huggingface/datasets/blob/9b21e181b642bd55b3ef68c1948bfbcd388136d6/src/datasets/info.py#L94)
ds1 = Dataset.from_generator(
generator=batch_hidden_states,
info=DatasetInfo(
description=json.dumps(info_kwargs, indent=2),
config_name=f,
),
gen_kwargs=gen_kwargs,
num_proc=1,
)
# ## Add labels
# For our probe. Given next_token scores (logits) we take only the subset the corresponds to our negative tokens (e.g. False, no, ...) and positive tokens (e.g. Yes, yes, affirmative, ...).
def expand_choices(choices: List[str]) -> List[str]:
"""expand out choices by adding versions that are upper, lower, whitespace, etc"""
new = []
for c in choices:
new.append(c)
new.append(c.upper())
new.append(c.capitalize())
new.append(c.lower())
return set(new)
# left_choices = list(r[0] for r in ds1['answer_choices'])+['no', 'false', 'negative', 'wrong']
# right_choices = list(r[1] for r in ds1['answer_choices'])+['yes', 'true', 'positive', 'right']
# left_choices, right_choices = expand_choices(left_choices), expand_choices(right_choices)
# assert len(set(left_choices).intersection(right_choices))==0
# expanded_choices = [left_choices, right_choices]
# expanded_choice_ids = choice2ids(expanded_choices, tokenizer)
# add_ans_exp = lambda r: scores2choice_probs(r, expanded_choice_ids, prefix="expanded_", keys=["scores0"])
# FIXME: we have the yes and no swapped
# FIXME: use k-mean closest tokens?
# this is just based on pairs for that answer...
# FIXME of course I added a dim!
add_txt_ans0 = lambda r: {'txt_ans0': tokenizer.batch_decode(torch.softmax(torch.tensor(r['scores0']), 0).argmax(0))}
# Either just use the template choices
add_ans = lambda r: scores2choice_probs(r, row_choice_ids(r, tokenizer), keys=["scores0"])
# Or all expanded choices
ds1.set_format(type='numpy')#, columns=['input_ids', 'token_type_ids', 'attention_mask', 'label'])
ds3 = (
ds1
.map(add_ans, desc='add_ans') # slow?
# .map(add_ans_exp)
.map(add_txt_ans0, desc='add_txt_ans0')
)
ds1 = create_hs_ds(ds_name, ds_tokens, model, cfg, intervention_dicts=intervention, f=f)
ds3 = post_proc_hs_ds(ds1, tokenizer)
ds3.save_to_disk(f)
print('! saved f=', f)
+7 -6
View File
@@ -6,7 +6,7 @@ from datasets.arrow_dataset import Dataset
import hashlib
import pickle
import numpy as np
from typing import List, Dict, Any, Union, NewType
from typing import List, Dict, Any, Union, NewType, Optional
from src.datasets.hs import ExtractHiddenStates
from src.helpers.typing import float_to_int16, int16_to_float
@@ -14,7 +14,7 @@ from src.helpers.ds import ds_keep_cols, clear_mem
from src.datasets.intervene import InterventionDict
def batch_hidden_states(model, tokenizer, intervention_dicts: List[InterventionDict], data: Dataset, batch_size=2, layer_padding=3, layer_stride=4):
def batch_hidden_states(model, tokenizer, intervention_dicts: Optional[InterventionDict], data: Dataset, batch_size=2, layer_padding=3, layer_stride=4):
"""
Given an encoder-decoder model, a list of data, computes the contrast hidden states on n random examples.
Returns numpy arrays of shape (n, hidden_dim) for each candidate label, along with a boolean numpy array of shape (n,)
@@ -47,21 +47,22 @@ def batch_hidden_states(model, tokenizer, intervention_dicts: List[InterventionD
large_arrays_keys = [k for k,v in hsl.items() if isinstance(v, torch.Tensor) and v.ndim>2]
# TODO deal with multiple lists of hs in hs0
large_arrays_as_int16 = {k:hsl[k][j] for k in large_arrays_keys}
large_arrays = {k:hsl[k][j] for k in large_arrays_keys}
yield dict(
# large_arrays_keys=large_arrays_keys,
scores0=hsl["scores"][j],
scores0=hsl["scores"][j],
# layer_names=hsl["layers"][j] if k==0 else [], # just in the first one, to save space
ds_index=index[j],
# int16 makes our storage much smaller
**large_arrays_as_int16,
**large_arrays,
**info
)
info = large_arrays_as_int16= hsl = None
info = large_arrays = hsl = None
clear_mem()
+32 -18
View File
@@ -31,7 +31,8 @@ from src.datasets.scores import choice2id, choice2ids
from src.helpers.torch import clear_mem
from collections import defaultdict
from dataclasses import field
from src.datasets.intervene import InterventionDict
from src.datasets.intervene import InterventionDict, intervention_meta_fn
from functools import partial
def noise_for_embeds(inputs_embeds, seed=42, std = 2e-2):
@@ -78,9 +79,22 @@ class ExtractHiddenStates:
model: PreTrainedModel
tokenizer: PreTrainedTokenizer
intervention_dicts: List[Optional[InterventionDict]]
intervention_dicts: Optional[InterventionDict] = None
layer_stride: int = 8
layer_padding: int = 3
def get_layer_names(self):
# for WizardLM/WizardCoder-3B-V1.0
# HEADS = [f"transformer.h.{i}.attn.c_proj" for i in range(self.model.config.num_hidden_layers)]
# MLPS = [f"transformer.h.{i}.mlp" for i in range(self.model.config.num_hidden_layers)]
# for "WizardLM/WizardCoder-Python-13B-V1.0"
# HACK: depends on model layout
layers_names = [f"model.layers.{i}.self_attn" for i in range(self.model.config.num_hidden_layers)]
module_names = [k for k,v in self.model.named_modules()]
layers_not_found = set(layers_names)-set(module_names)
assert len(layers_not_found)==0, f"some layers not found in model: {layers_not_found}. we have {layers_names}"
return self.get_layer_selection(layers_names)
def get_batch_of_hidden_states(
@@ -117,26 +131,26 @@ class ExtractHiddenStates:
# forward pass
last_token = -1
# for WizardLM/WizardCoder-3B-V1.0
# HEADS = [f"transformer.h.{i}.attn.c_proj" for i in range(self.model.config.num_hidden_layers)]
# MLPS = [f"transformer.h.{i}.mlp" for i in range(self.model.config.num_hidden_layers)]
# for "WizardLM/WizardCoder-Python-13B-V1.0"
# HACK: depends on model layout
layers_names = [f"model.layers.{i}.self_attn" for i in range(self.model.config.num_hidden_layers)]
# MLPS = [f"model.layers.{i}.mlp" for i in range(self.model.config.num_hidden_layers)]
module_names = [k for k,v in self.model.named_modules()]
layers_not_found = set(layers_names)-set(module_names)
assert len(layers_not_found)==0, f"some layers not found in model: {layers_not_found}. we have {layers_names}"
layers_names, layer_inds = self.get_layer_selection(layers_names)
layers_names, layer_inds = self.get_layer_names()
self.model.eval()
# outs = []
if self.intervention_dicts is not None:
# extraction mode
# 15 is a magic number from honest_llama
num_heads = self.model.config.num_attention_heads
intervention_fn1 = partial(intervention_meta_fn, interventions=self.intervention_dicts, num_heads=num_heads, alpha=-15)
intervention_fn2 = partial(intervention_meta_fn, interventions=self.intervention_dicts, num_heads=num_heads, alpha=15)
edit_outputs = [intervention_fn1, intervention_fn2]
else:
# calibration mode
edit_outputs = [None]
with torch.no_grad():
multi_outs = defaultdict(list)
for edit_output in self.intervention_dicts:
for edit_output in edit_outputs:
with TraceDict(self.model, layers_names, retain_grad=True, detach=True, edit_output=edit_output) as ret:
model_inputs = self.model.prepare_inputs_for_generation(input_ids=input_ids, attention_mask=attention_mask, use_cache=False)
outputs = self.model.forward(
+13 -1
View File
@@ -175,9 +175,21 @@ def load_prompts(
prompt_format=prompt_format,
)
prompts = [{'ds_string': ds_string, 'example_i':i, **p} for p in prompts]
def prompt_ok(prompt):
""" we want answers where we can distinguish them from the first token
we don't have access to the tokenizer here, so we just make sure the first 3 letters are differen't and there are not spaces
"""
answer_choices = prompt['answer_choices']
a = answer_choices[0][:3]
b = answer_choices[1][:3]
return (a != b) and ' ' not in a
prompts = list(filter(prompt_ok, prompts))
prompts = prompt_sampler(prompts, seed=42+j)
# TODO: make sure they are single token answers (or at least the first token is unique)
for p in prompts:
j +=1
j += 1
yield p