mirror of
https://github.com/wassname/discovering_latent_knowledge.git
synced 2026-09-26 13:40:30 +08:00
🧠💭
This commit is contained in:
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user