mirror of
https://github.com/wassname/LoRA_are_lie_detectors.git
synced 2026-08-20 12:00:58 +08:00
wip collect residuals from adapter layers
This commit is contained in:
@@ -1803,7 +1803,9 @@
|
||||
"source": [
|
||||
"# # for single process DEBUGING\n",
|
||||
"# from src.eval.collect import generate_batches\n",
|
||||
"# o = next(iter(generate_batches(dl_OOD, model)))\n"
|
||||
"# o = next(iter(generate_batches(dl_OOD, model)))\n",
|
||||
"\n",
|
||||
"collection_layers = cfg.collection_layers\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -1858,8 +1860,8 @@
|
||||
],
|
||||
"source": [
|
||||
"dataset_dir=Path(trainer.log_dir)/'hidden_states'\n",
|
||||
"ds_out_OOD, f = manual_collect2(dl_OOD, model, dataset_name=\"OOD\", dataset_dir=dataset_dir)\n",
|
||||
"ds_out_valtest, f = manual_collect2(dl_valtest2, model, dataset_name=\"valtest\", dataset_dir=dataset_dir)\n"
|
||||
"ds_out_OOD, f = manual_collect2(dl_OOD, model, dataset_name=\"OOD\", layers=collection_layers, dataset_dir=dataset_dir)\n",
|
||||
"ds_out_valtest, f = manual_collect2(dl_valtest2, model, dataset_name=\"valtest\", layers=collection_layers, dataset_dir=dataset_dir)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -2107,6 +2109,14 @@
|
||||
"File \u001b[0;32m/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/src/eval/interventions.py:11\u001b[0m\n\u001b[1;32m 9\u001b[0m \u001b[39mfrom\u001b[39;00m \u001b[39msklearn\u001b[39;00m\u001b[39m.\u001b[39;00m\u001b[39mmodel_selection\u001b[39;00m \u001b[39mimport\u001b[39;00m train_test_split\n\u001b[1;32m 10\u001b[0m \u001b[39mfrom\u001b[39;00m \u001b[39msrc\u001b[39;00m\u001b[39m.\u001b[39;00m\u001b[39mprobes\u001b[39;00m\u001b[39m.\u001b[39;00m\u001b[39mutils\u001b[39;00m \u001b[39mimport\u001b[39;00m postproc, make_dfres_pretty\n\u001b[0;32m---> 11\u001b[0m \u001b[39mfrom\u001b[39;00m \u001b[39msrc\u001b[39;00m\u001b[39m.\u001b[39;00m\u001b[39mprobes\u001b[39;00m\u001b[39m.\u001b[39;00m\u001b[39msk_lr\u001b[39;00m \u001b[39mimport\u001b[39;00m check_lr_intervention_predictive\n\u001b[1;32m 14\u001b[0m \u001b[39mdef\u001b[39;00m \u001b[39mtest_intervention_quality2\u001b[39m(ds_out, label_fn, thresh\u001b[39m=\u001b[39m\u001b[39m0.03\u001b[39m, take_diff\u001b[39m=\u001b[39m\u001b[39mFalse\u001b[39;00m, verbose\u001b[39m=\u001b[39m\u001b[39mFalse\u001b[39;00m, title\u001b[39m=\u001b[39m\u001b[39m\"\u001b[39m\u001b[39mIntervention predictive power\u001b[39m\u001b[39m\"\u001b[39m, skip\u001b[39m=\u001b[39m\u001b[39m0\u001b[39m, stride\u001b[39m=\u001b[39m\u001b[39m1\u001b[39m, model_kwargs\u001b[39m=\u001b[39m{}):\n\u001b[1;32m 15\u001b[0m \u001b[39m \u001b[39m\u001b[39m\"\"\"\u001b[39;00m\n\u001b[1;32m 16\u001b[0m \u001b[39m Check interventions are ordered and different and valid\u001b[39;00m\n\u001b[1;32m 17\u001b[0m \n\u001b[0;32m (...)\u001b[0m\n\u001b[1;32m 24\u001b[0m \u001b[39m - it's not over confident\u001b[39;00m\n\u001b[1;32m 25\u001b[0m \u001b[39m \"\"\"\u001b[39;00m\n",
|
||||
"\u001b[0;31mModuleNotFoundError\u001b[0m: No module named 'src.probes.sk_lr'"
|
||||
]
|
||||
},
|
||||
{
|
||||
"ename": "",
|
||||
"evalue": "",
|
||||
"output_type": "error",
|
||||
"traceback": [
|
||||
"\u001b[1;31mThe Kernel crashed while executing code in the the current cell or a previous cell. Please review the code in the cell(s) to identify a possible cause of the failure. Click <a href='https://aka.ms/vscodeJupyterKernelCrash'>here</a> for more info. View Jupyter <a href='command:jupyter.viewOutput'>log</a> for further details."
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -492,7 +492,8 @@
|
||||
" def __init__(\n",
|
||||
" self,\n",
|
||||
" c_in,\n",
|
||||
" total_steps,\n",
|
||||
" epoch_steps,\n",
|
||||
" max_epochs,\n",
|
||||
" depth=0,\n",
|
||||
" lr=4e-3,\n",
|
||||
" weight_decay=1e-9,\n",
|
||||
@@ -502,7 +503,7 @@
|
||||
" dropout=0,\n",
|
||||
" **kwargs,\n",
|
||||
" ):\n",
|
||||
" super().__init__(total_steps=total_steps, lr=lr, weight_decay=weight_decay)\n",
|
||||
" super().__init__(epoch_steps=epoch_steps, max_epochs=max_epochs, lr=lr, weight_decay=weight_decay)\n",
|
||||
" self.save_hyperparameters()\n",
|
||||
"\n",
|
||||
" self.ae = AutoEncoder(\n",
|
||||
@@ -769,13 +770,22 @@
|
||||
"\u001b[1;32m/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/notebooks/11_vae_w_importance.ipynb Cell 20\u001b[0m line \u001b[0;36m1\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/notebooks/11_vae_w_importance.ipynb#X14sZmlsZQ%3D%3D?line=5'>6</a>\u001b[0m \u001b[39mdef\u001b[39;00m \u001b[39m__init__\u001b[39m(\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/notebooks/11_vae_w_importance.ipynb#X14sZmlsZQ%3D%3D?line=6'>7</a>\u001b[0m \u001b[39mself\u001b[39m,\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/notebooks/11_vae_w_importance.ipynb#X14sZmlsZQ%3D%3D?line=7'>8</a>\u001b[0m c_in,\n\u001b[0;32m (...)\u001b[0m\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/notebooks/11_vae_w_importance.ipynb#X14sZmlsZQ%3D%3D?line=16'>17</a>\u001b[0m \u001b[39m*\u001b[39m\u001b[39m*\u001b[39mkwargs,\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/notebooks/11_vae_w_importance.ipynb#X14sZmlsZQ%3D%3D?line=17'>18</a>\u001b[0m ):\n\u001b[0;32m---> <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/notebooks/11_vae_w_importance.ipynb#X14sZmlsZQ%3D%3D?line=18'>19</a>\u001b[0m \u001b[39msuper\u001b[39;49m()\u001b[39m.\u001b[39;49m\u001b[39m__init__\u001b[39;49m(total_steps\u001b[39m=\u001b[39;49mtotal_steps, lr\u001b[39m=\u001b[39;49mlr, weight_decay\u001b[39m=\u001b[39;49mweight_decay)\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/notebooks/11_vae_w_importance.ipynb#X14sZmlsZQ%3D%3D?line=19'>20</a>\u001b[0m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39msave_hyperparameters()\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/notebooks/11_vae_w_importance.ipynb#X14sZmlsZQ%3D%3D?line=21'>22</a>\u001b[0m \u001b[39mself\u001b[39m\u001b[39m.\u001b[39mae \u001b[39m=\u001b[39m AutoEncoder(\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/notebooks/11_vae_w_importance.ipynb#X14sZmlsZQ%3D%3D?line=22'>23</a>\u001b[0m c_in,\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/notebooks/11_vae_w_importance.ipynb#X14sZmlsZQ%3D%3D?line=23'>24</a>\u001b[0m n_hidden\u001b[39m=\u001b[39mhs,\n\u001b[0;32m (...)\u001b[0m\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/notebooks/11_vae_w_importance.ipynb#X14sZmlsZQ%3D%3D?line=27'>28</a>\u001b[0m dropout\u001b[39m=\u001b[39mdropout,\n\u001b[1;32m <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/notebooks/11_vae_w_importance.ipynb#X14sZmlsZQ%3D%3D?line=28'>29</a>\u001b[0m )\n",
|
||||
"\u001b[0;31mTypeError\u001b[0m: PLBase.__init__() got an unexpected keyword argument 'total_steps'"
|
||||
]
|
||||
},
|
||||
{
|
||||
"ename": "",
|
||||
"evalue": "",
|
||||
"output_type": "error",
|
||||
"traceback": [
|
||||
"\u001b[1;31mThe Kernel crashed while executing code in the the current cell or a previous cell. Please review the code in the cell(s) to identify a possible cause of the failure. Click <a href='https://aka.ms/vscodeJupyterKernelCrash'>here</a> for more info. View Jupyter <a href='command:jupyter.viewOutput'>log</a> for further details."
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"\n",
|
||||
"net = PLAE(\n",
|
||||
" c_in=c_in,\n",
|
||||
" total_steps=max_epochs * len(dl_train) * VAE_EPOCH_MULT,\n",
|
||||
" epoch_steps=max_epochs,\n",
|
||||
" max_epochs=max_epochs * VAE_EPOCH_MULT,\n",
|
||||
" lr=lr,\n",
|
||||
" weight_decay=wd,\n",
|
||||
" hs=32,\n",
|
||||
|
||||
Generated
+59
-1
@@ -198,6 +198,26 @@ docs = ["furo", "myst-parser", "sphinx", "sphinx-notfound-page", "sphinxcontrib-
|
||||
tests = ["attrs[tests-no-zope]", "zope-interface"]
|
||||
tests-no-zope = ["cloudpickle", "hypothesis", "mypy (>=1.1.1)", "pympler", "pytest (>=4.3.0)", "pytest-mypy-plugins", "pytest-xdist[psutil]"]
|
||||
|
||||
[[package]]
|
||||
name = "baukit"
|
||||
version = "0.0.1"
|
||||
description = ""
|
||||
optional = false
|
||||
python-versions = ">=3.7"
|
||||
files = []
|
||||
develop = false
|
||||
|
||||
[package.dependencies]
|
||||
numpy = "*"
|
||||
torch = "*"
|
||||
torchvision = "*"
|
||||
|
||||
[package.source]
|
||||
type = "git"
|
||||
url = "https://github.com/davidbau/baukit"
|
||||
reference = "HEAD"
|
||||
resolved_reference = "5e23007c02fd58f063200c5dc9033e90f092630d"
|
||||
|
||||
[[package]]
|
||||
name = "bitsandbytes"
|
||||
version = "0.41.3.post2"
|
||||
@@ -3455,6 +3475,44 @@ text = ["nltk (>=3.6)", "regex (>=2021.9.24)", "tqdm (>=4.41.0)", "transformers
|
||||
typing = ["mypy (==1.7.1)", "torch (==2.1.1)", "types-PyYAML", "types-emoji", "types-protobuf", "types-requests", "types-setuptools", "types-six", "types-tabulate"]
|
||||
visual = ["SciencePlots (>=2.0.0)", "matplotlib (>=3.2.0)"]
|
||||
|
||||
[[package]]
|
||||
name = "torchvision"
|
||||
version = "0.16.2"
|
||||
description = "image and video datasets and models for torch deep learning"
|
||||
optional = false
|
||||
python-versions = ">=3.8"
|
||||
files = [
|
||||
{file = "torchvision-0.16.2-cp310-cp310-macosx_10_13_x86_64.whl", hash = "sha256:bc86f2800cb2c0c1a09c581409cdd6bff66e62f103dc83fc63f73346264c3756"},
|
||||
{file = "torchvision-0.16.2-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:b024bd412df6d3a007dcebf311a894eb3c5c21e1af80d12be382bbcb097a7c3a"},
|
||||
{file = "torchvision-0.16.2-cp310-cp310-manylinux1_x86_64.whl", hash = "sha256:e89f10f3c8351972b6e3fda95bc3e479ea8dbfc9dfcfd2c32902dbad4ba5cfc5"},
|
||||
{file = "torchvision-0.16.2-cp310-cp310-manylinux2014_aarch64.whl", hash = "sha256:96c7583700112a410bdc4e1e4f118c429dab49c29c9a31a2cc3579bc9b08b19d"},
|
||||
{file = "torchvision-0.16.2-cp310-cp310-win_amd64.whl", hash = "sha256:9f4032ebb3277fb07ff6a9b818d50a547fb8fcd89d958cfd9e773322454bb688"},
|
||||
{file = "torchvision-0.16.2-cp311-cp311-macosx_10_13_x86_64.whl", hash = "sha256:67b1aaf8b8cb02ce75dd445f291a27c8036a502f8c0aa76e28c37a0faac2e153"},
|
||||
{file = "torchvision-0.16.2-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:bef30d03e1d1c629761f4dca51d3b7d8a0dc0acce6f4068ab2a1634e8e7b64e0"},
|
||||
{file = "torchvision-0.16.2-cp311-cp311-manylinux1_x86_64.whl", hash = "sha256:e59cc7b2bd1ab5c0ce4ae382e4e37be8f1c174e8b5de2f6a23c170de9ae28495"},
|
||||
{file = "torchvision-0.16.2-cp311-cp311-manylinux2014_aarch64.whl", hash = "sha256:e130b08cc9b3cc73a6c59d6edf032394a322f9579bfd21d14bc2e1d0999aa758"},
|
||||
{file = "torchvision-0.16.2-cp311-cp311-win_amd64.whl", hash = "sha256:8692ab1e48807e9604046a6f4beeb67b523294cee1b00828654bb0df2cfce2b2"},
|
||||
{file = "torchvision-0.16.2-cp38-cp38-macosx_10_13_x86_64.whl", hash = "sha256:b82732dcf876a37c852772342aa6ee3480c03bb3e2a802ae109fc5f7e28d26e9"},
|
||||
{file = "torchvision-0.16.2-cp38-cp38-macosx_11_0_arm64.whl", hash = "sha256:4b065143d1a720fe8a9077fd4be35d491f98819ec80b3dbbc3ec64d0b707a906"},
|
||||
{file = "torchvision-0.16.2-cp38-cp38-manylinux1_x86_64.whl", hash = "sha256:bc5f274e4ecd1b86062063cdf4fd385a1d39d147a3a2685fbbde9ff08bb720b8"},
|
||||
{file = "torchvision-0.16.2-cp38-cp38-manylinux2014_aarch64.whl", hash = "sha256:335959c43b371c0474af34c1ef2a52efdc7603c45700d29e4475eeb02984170c"},
|
||||
{file = "torchvision-0.16.2-cp38-cp38-win_amd64.whl", hash = "sha256:7fd22d86e08eba321af70cad291020c2cdeac069b00ce88b923ca52e06174769"},
|
||||
{file = "torchvision-0.16.2-cp39-cp39-macosx_10_13_x86_64.whl", hash = "sha256:56115268b37f0b75364e3654e47ad9abc66ac34c1f9e5e3dfa89a22d6a40017a"},
|
||||
{file = "torchvision-0.16.2-cp39-cp39-macosx_11_0_arm64.whl", hash = "sha256:82805f8445b094f9d1e770390ee6cc86855e89955e08ce34af2e2274fc0e5c45"},
|
||||
{file = "torchvision-0.16.2-cp39-cp39-manylinux1_x86_64.whl", hash = "sha256:3f4bd5fcbc361476e2e78016636ac7d5509e59d9962521f06eb98e6803898182"},
|
||||
{file = "torchvision-0.16.2-cp39-cp39-manylinux2014_aarch64.whl", hash = "sha256:8199acdf8ab066a28b84a5b6f4d97b58976d9e164b1acc3a9d14fccfaf74bb3a"},
|
||||
{file = "torchvision-0.16.2-cp39-cp39-win_amd64.whl", hash = "sha256:41dd4fa9f176d563fe9f1b9adef3b7e582cdfb60ce8c9bc51b094a025be687c9"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
numpy = "*"
|
||||
pillow = ">=5.3.0,<8.3.dev0 || >=8.4.dev0"
|
||||
requests = "*"
|
||||
torch = "2.1.2"
|
||||
|
||||
[package.extras]
|
||||
scipy = ["scipy"]
|
||||
|
||||
[[package]]
|
||||
name = "tornado"
|
||||
version = "6.4"
|
||||
@@ -3918,4 +3976,4 @@ multidict = ">=4.0"
|
||||
[metadata]
|
||||
lock-version = "2.0"
|
||||
python-versions = ">=3.10,<3.13"
|
||||
content-hash = "89138036efeb67905cec4ab43b2672dad018fad0503393c79ae3b4457062e279"
|
||||
content-hash = "7434f6a628f6d815499a2240abbf2b1942e30b4373f85d53ef479ed61a5ead08"
|
||||
|
||||
@@ -33,6 +33,7 @@ ipywidgets = "^8.1.1"
|
||||
tabulate = "^0.9.0"
|
||||
seaborn = "^0.13.0"
|
||||
mapie = "^0.7.0"
|
||||
baukit = {git = "https://github.com/davidbau/baukit"}
|
||||
|
||||
[[tool.poetry.source]]
|
||||
name = "pytorch"
|
||||
|
||||
@@ -454,6 +454,15 @@ Logically what's important are the activations on the default params! So I shoul
|
||||
# for normal ia3 seetings see https://github.com/huggingface/peft/blob/cf04d0353f0343cbf66627228c4495f51669af34/src/peft/utils/constants.py#L81
|
||||
# and https://github.com/huggingface/peft/blob/cf04d0353f0343cbf66627228c4495f51669af34/src/peft/utils/constants.py#L102
|
||||
|
||||
|
||||
"help": (
|
||||
"List of module names or regex expression of the module names to replace with LoRA."
|
||||
"For example, ['q', 'v'] or '.*decoder.*(SelfAttention|EncDecAttention).*(q|v)$'. "
|
||||
"If not specified, modules will be chosen according to the model architecture, If the architecture is "
|
||||
"not known, an error will be raised -- in this case, you shoud specify the target modules manually."
|
||||
),
|
||||
|
||||
|
||||
TODO:
|
||||
- [ ] read anthropic [paper](https://transformer-circuits.pub/2022/toy_model/index.html) on importance matrix,
|
||||
- [x] [maybe reply to colin](https://www.lesswrong.com/posts/LnHowHgmrMbWtpkxx/intro-to-superposition-and-sparse-autoencoders-colab)
|
||||
|
||||
@@ -21,6 +21,9 @@ class ExtractConfig(Serializable):
|
||||
# model: str = "/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/phi-1_5"
|
||||
model: str = "wassname/phi-1_5-w_hidden_states"
|
||||
|
||||
# collection_layers: tuple[str, ...] = ("layer.0", "layer.1", "layer.2", "layer.3", "layer.4", "layer.5", "layer.6", "layer.7", "layer.8", "layer.9", "layer.10", "layer.11")
|
||||
# """Names of layers to extract from using baukit.nethook.TraceDict"""
|
||||
|
||||
batch_size: int = 2
|
||||
|
||||
prompt_format: str | None = 'phi'
|
||||
|
||||
+15
-9
@@ -6,12 +6,15 @@ from tqdm.auto import tqdm
|
||||
from torch.utils.data import random_split, DataLoader, TensorDataset
|
||||
import json
|
||||
from loguru import logger
|
||||
|
||||
from baukit.nethook import TraceDict
|
||||
|
||||
from src.helpers.torch_helpers import clear_mem, detachcpu, recursive_copy
|
||||
from src.models.pl_lora_ft import postprocess_result
|
||||
from src.config import root_folder
|
||||
|
||||
@torch.no_grad
|
||||
def generate_batches(loader: DataLoader, model: AutoModelForCausalLM, get_residual=True) -> dict:
|
||||
def generate_batches(loader: DataLoader, model: AutoModelForCausalLM, layers, get_residual=True) -> dict:
|
||||
if not hasattr(model, 'disable_adapter'):
|
||||
logger.warning("model does not have disable_adapter")
|
||||
model.eval()
|
||||
@@ -23,15 +26,18 @@ def generate_batches(loader: DataLoader, model: AutoModelForCausalLM, get_residu
|
||||
)
|
||||
if hasattr(model, 'disable_adapter'):
|
||||
with model.disable_adapter():
|
||||
out = model(**b_in, use_cache=False, output_hidden_states=True, return_dict=True)
|
||||
res = {f'{k}_base':v for k,v in postprocess_result(batch, out, get_residual=get_residual).items()}
|
||||
with TraceDict(model, layers, detach=True) as ret:
|
||||
out = model(**b_in, use_cache=False, output_hidden_states=True, return_dict=True)
|
||||
res = {f'{k}_base':v for k,v in postprocess_result(batch, ret, out, get_residual=get_residual).items()}
|
||||
del out
|
||||
out_a = model(**b_in, use_cache=False, output_hidden_states=True, return_dict=True)
|
||||
res_a = {f'{k}_adapt':v for k,v in postprocess_result(batch, out_a, get_residual=get_residual).items()}
|
||||
with TraceDict(model, layers, detach=True) as ret_a:
|
||||
out_a = model(**b_in, use_cache=False, output_hidden_states=True, return_dict=True)
|
||||
res_a = {f'{k}_adapt':v for k,v in postprocess_result(batch, ret_a, out_a, get_residual=get_residual).items()}
|
||||
del out_a
|
||||
else:
|
||||
out = model(**b_in, use_cache=False, output_hidden_states=True, return_dict=True)
|
||||
res = {f'{k}_base':v for k,v in postprocess_result(batch, out, get_residual=get_residual).items()}
|
||||
with TraceDict(model, layers) as ret:
|
||||
out = model(**b_in, use_cache=False, output_hidden_states=True, return_dict=True)
|
||||
res = {f'{k}_base':v for k,v in postprocess_result(batch, ret, out,get_residual=get_residual).items()}
|
||||
res_a = {}
|
||||
|
||||
o = dict(**res, **res_a)
|
||||
@@ -49,13 +55,13 @@ def ds_hash(**kwargs):
|
||||
return suffix
|
||||
|
||||
|
||||
def manual_collect2(loader: DataLoader, model: AutoModelForCausalLM, dataset_name='', get_residual=True, dataset_dir=root_folder):
|
||||
def manual_collect2(loader: DataLoader, model: AutoModelForCausalLM, dataset_name='', layers=[], get_residual=True, dataset_dir=root_folder):
|
||||
hash = ds_hash(generate_batches=generate_batches, loader=loader, model=model)
|
||||
f = dataset_dir / ".ds" / f"ds_{dataset_name}_{hash}"
|
||||
f.parent.mkdir(exist_ok=True, parents=True)
|
||||
f = str(f)
|
||||
logger.info(f"creating dataset {f}")
|
||||
iterator = generate_batches(loader, model, get_residual=get_residual)
|
||||
iterator = generate_batches(loader, model, layers=layers, get_residual=get_residual)
|
||||
with ArrowWriter(path=f, writer_batch_size=6) as writer:
|
||||
for bo in iterator:
|
||||
# dict_of_batches_to_batch_of_dicts
|
||||
|
||||
@@ -1,44 +1,6 @@
|
||||
import torch
|
||||
# import numpy as np
|
||||
# import transformers
|
||||
# import random
|
||||
import gc
|
||||
# import pandas as pd
|
||||
|
||||
# def get_top_n(scores: torch.Tensor, tokenizer: transformers.PreTrainedTokenizer, n=10) -> pd.Series:
|
||||
# """Get top n choices and their probabilities given raw logits"""
|
||||
# probs = scores.softmax(-1).squeeze()
|
||||
# assert len(probs.shape)==1
|
||||
# top10 = torch.argsort(probs, dim=-1, descending=True)[:n]
|
||||
# top10_probs = probs[top10]
|
||||
# top10_ext = tokenizer.batch_decode(top10)
|
||||
# return pd.Series(top10_probs, index=top10_ext, name='probs')
|
||||
|
||||
# def to_numpy(x):
|
||||
# """
|
||||
# Trys to convert torch to numpy and if possible a single item
|
||||
# """
|
||||
# if isinstance(x, torch.Tensor):
|
||||
# # note apache parquet doesn't support half https://github.com/huggingface/datasets/issues/4981
|
||||
# x = x.detach().cpu().float()
|
||||
# if x.squeeze().dim()==0:
|
||||
# return x.item()
|
||||
# return x.numpy()
|
||||
# else:
|
||||
# return x
|
||||
|
||||
|
||||
|
||||
# def set_seeds(n: int) -> None:
|
||||
# transformers.set_seed(n)
|
||||
# torch.manual_seed(n)
|
||||
# np.random.seed(n)
|
||||
# random.seed(n)
|
||||
|
||||
# def to_item(x):
|
||||
# if isinstance(x, torch.Tensor):
|
||||
# x = x.detach().cpu().item()
|
||||
# return x
|
||||
import copy
|
||||
|
||||
from jaxtyping import Float, Int
|
||||
from torch import Tensor
|
||||
@@ -61,10 +23,7 @@ def detachcpu(x):
|
||||
Trys to convert torch if possible a single item
|
||||
"""
|
||||
if isinstance(x, torch.Tensor):
|
||||
# note apache parquet doesn't support half to we go for float https://github.com/huggingface/datasets/issues/4981
|
||||
x = x.detach().cpu()
|
||||
# if x.squeeze().dim()==0:
|
||||
# return x.item()
|
||||
x = x.cpu()
|
||||
return x
|
||||
else:
|
||||
return x
|
||||
@@ -91,11 +50,11 @@ def recursive_copy(x, clone=None, detach=None, retain_grad=None):
|
||||
return x
|
||||
# Only dicts, lists, and tuples (and subclasses) can be copied.
|
||||
if isinstance(x, dict):
|
||||
return type(x)({k: recursive_copy(v) for k, v in x.items()})
|
||||
return type(x)({k: recursive_copy(v, clone=clone, detach=detach, retain_grad=retain_grad) for k, v in x.items()})
|
||||
elif isinstance(x, (list, tuple)):
|
||||
return type(x)([recursive_copy(v) for v in x])
|
||||
return type(x)([recursive_copy(v, clone=clone, detach=detach, retain_grad=retain_grad) for v in x])
|
||||
else:
|
||||
assert False, f"Unknown type {type(x)} cannot be broken into tensors."
|
||||
return copy.deepcopy(x)
|
||||
|
||||
def batch_to_device(b, device=None):
|
||||
"""Move a batch to the device"""
|
||||
|
||||
+37
-17
@@ -8,6 +8,8 @@ from einops import rearrange
|
||||
from transformers.modeling_outputs import ModelOutput
|
||||
from jaxtyping import Float, Int
|
||||
from torch import Tensor
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
from baukit.nethook import TraceDict
|
||||
|
||||
from src.helpers.torch_helpers import clear_mem, detachcpu, recursive_copy, switch
|
||||
|
||||
@@ -19,17 +21,17 @@ def hacky_sanitize_outputs(o):
|
||||
return o
|
||||
|
||||
|
||||
def postprocess_result(i, o, get_residual=True):
|
||||
def postprocess_result(input: dict, ret: TraceDict, output: ModelOutput, get_residual=True) -> ModelOutput:
|
||||
|
||||
# note that the results are huge. It might be worth convertting to int16 or similar so we can save to disc as we go https://github.com/EleutherAI/elk/blob/84e99a36a5050881d85f1510a2486ce46ac1f942/elk/utils/typing.py#L16
|
||||
assert torch.isfinite(o['logits']).all()
|
||||
assert torch.isfinite(output['logits']).all()
|
||||
|
||||
|
||||
end_logits = o["logits"][:, -1].detach().cpu().float()
|
||||
end_logits = output["logits"][:, -1].detach().cpu().float()
|
||||
probs = torch.softmax(end_logits, -1)
|
||||
choice_ids = i['choice_ids'].detach().cpu().long()
|
||||
choice_ids = input['choice_ids'].detach().cpu().long()
|
||||
|
||||
label_instructed = i['label_true'] ^ i['instructed_to_lie']
|
||||
label_instructed = input['label_true'] ^ input['instructed_to_lie']
|
||||
|
||||
|
||||
choice_probs = select_choices(probs, choice_ids).sum(2)
|
||||
@@ -38,7 +40,7 @@ def postprocess_result(i, o, get_residual=True):
|
||||
binary_ans = choice_probs[:, 1] / (choice_probs.sum(1) + 1e-12)
|
||||
|
||||
|
||||
correct_truth_telling = switch(binary_ans, i['label_true'])
|
||||
correct_truth_telling = switch(binary_ans, input['label_true'])
|
||||
correct_instruction_following = switch(binary_ans, label_instructed)
|
||||
|
||||
out = dict(
|
||||
@@ -47,22 +49,40 @@ def postprocess_result(i, o, get_residual=True):
|
||||
# maybe these ones should be postprocessing
|
||||
choice_probs=choice_probs,
|
||||
binary_ans=binary_ans,
|
||||
label_true=i['label_true'],
|
||||
label_true=input['label_true'],
|
||||
label_instructed=label_instructed,
|
||||
instructed_to_lie=i['instructed_to_lie'],
|
||||
sys_instr_name=i['sys_instr_name'],
|
||||
example_i=i['example_i'],
|
||||
ds_string=i['ds_string'],
|
||||
template_name=i['template_name'],
|
||||
instructed_to_lie=input['instructed_to_lie'],
|
||||
sys_instr_name=input['sys_instr_name'],
|
||||
example_i=input['example_i'],
|
||||
ds_string=input['ds_string'],
|
||||
template_name=input['template_name'],
|
||||
correct_truth_telling=correct_truth_telling,
|
||||
correct_instruction_following=correct_instruction_following,
|
||||
)
|
||||
if get_residual:
|
||||
# hidden states come at as lists of layers, lets stack them
|
||||
hidden_states = rearrange(list(o['hidden_states']), 'l b t h -> b l t h').detach().cpu().float()
|
||||
end_hidden_states = hidden_states[:, :, -1, :]
|
||||
end_residual_stream = end_hidden_states.diff(1)
|
||||
out['end_residual_stream'] = end_residual_stream
|
||||
# we can also get activations from layers monitored in baukit
|
||||
activations = {}
|
||||
for k in ret.keys():
|
||||
suffix = k.split('.')[-1]
|
||||
if suffix not in activations:
|
||||
activations[suffix] = []
|
||||
activations[suffix].append(ret[k].output)
|
||||
|
||||
for k in activations.keys():
|
||||
# HACK: we will assume they are all shaped [batch, tokens, hidden]
|
||||
activation = rearrange(activations[k], 'l b t h -> b l t h').detach().cpu().float()
|
||||
end_activation = activation[:, :, -1, :]
|
||||
end_residual = end_activation.diff(1)
|
||||
out[f'end_residual_{k}'] = end_residual
|
||||
|
||||
# ret = {k: v.detach().cpu().float() for k, v in ret.items()}
|
||||
|
||||
|
||||
# # hidden states come at as lists of layers, lets stack them
|
||||
# hidden_states = rearrange(list(output['hidden_states']), 'l b t h -> b l t h').detach().cpu().float()
|
||||
# end_hidden_states = hidden_states[:, :, -1, :]
|
||||
# end_residual_stream = end_hidden_states.diff(1)
|
||||
# out['end_residual_stream'] = end_residual_stream
|
||||
|
||||
# why oh why do I get mem leaks like this
|
||||
out = hacky_sanitize_outputs(out)
|
||||
|
||||
Reference in New Issue
Block a user