mirror of
https://github.com/wassname/LoRA_are_lie_detectors.git
synced 2026-09-12 12:04:56 +08:00
21 KiB
21 KiB
In [ ]:
import os
import numpy as np
import pandas as pd
from matplotlib import pyplot as plt
from tqdm.auto import tqdm
from typing import Optional, List, Dict, Union
from jaxtyping import Float
from torch import Tensor
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import Tensor
from torch import optim
from torch.utils.data import random_split, DataLoader, TensorDataset
from pathlib import Path
from einops import rearrange
import transformers
from transformers import (
AutoTokenizer,
AutoModelForCausalLM,
BitsAndBytesConfig,
AutoConfig,
)
from peft import (
get_peft_config,
get_peft_model,
LoraConfig,
TaskType,
LoftQConfig,
IA3Config,
)
from pathlib import Path
import datasets
from datasets import Dataset
from loguru import logger
logger.add(os.sys.stderr, format="{time} {level} {message}", level="INFO")
# load my code
%load_ext autoreload
%autoreload 2
import lightning.pytorch as pl
from src.config import ExtractConfig
from src.llms.load import load_model
from src.helpers.torch_helpers import clear_mem
from src.llms.phi.model_phi import PhiForCausalLMWHS
from src.eval.ds import filter_ds_to_known
from src.datasets.act_dm import ActivationDataModule
# plt.style.use("ggplot")
# plt.style.use("seaborn-v0_8")
import seaborn as sns
sns.set_theme('paper')
In [ ]:
# params
# cfg = ExtractConfig(
# # model="microsoft/phi-2",
# # # batch_size=1,
# # prompt_format="phi",
# )
# cfg
# params
batch_size = 32
lr = 4e-3
wd = 1e-4
MAX_ROWS = 2000
SKIP=5 # skip initial N layers
STRIDE=4 # skip every N layers
DECIMATE=2 # discard N features for speed
device = "cuda:0"
max_epochs = 44
VAE_EPOCH_MULT = 5
l1_coeff = 1.0e-1 # neel uses 3e-4 ! https://github.dev/neelnanda-io/1L-Sparse-Autoencoder/blob/bcae01328a2f41d24bd4a9160828f2fc22737f75/utils.py#L106, but them they sum l1 where mean l2
# x_feats=x_feats. other use 1e-1
BASE_FOLDER = Path("/media/wassname/SGIronWolf/projects5/elk/sgd_probes_are_lie_detectors/notebooks/lightning_logs/version_24/")
layers_names = ('fc1', 'Wqkv', 'fc2', 'out_proj')In [ ]:
# load hidden state from a previously loaded adapter
# the columns with _base are from the base model, and adapt from adapter
# FROM TRAINING TRUTH
f1_val = next(iter(BASE_FOLDER.glob('hidden_states/.ds/ds_valtest_*')))
f1_ood = next(iter(BASE_FOLDER.glob('hidden_states/.ds/ds_OOD_*')))
f1_val, f1_oodIn [ ]:
# insample_datasets = list(set(ds_val['ds_string_base']))
# outsample_datasets = list(set(ds_ood['ds_string_base']))
# print(insample_datasets, outsample_datasets)In [ ]:
input_columns = ['binary_ans_base', 'binary_ans_adapt' ] + [f'end_residual_{layer}_base' for layer in layers_names] + [f'end_residual_{layer}_adapt' for layer in layers_names]
def ds2xy_batched(ds):
data = []
for layer in layers_names:
# Stack the base and adapter representations as a 4th dim
X1 = [ds[f'end_residual_{layer}_base'], ds[f'end_residual_{layer}_adapt']]
X1 = rearrange(X1, 'versions b l f -> b l f versions')
data.append(X1)
# concat layers
# x = rearrange(data, 'b parts l f v -> b l (parts f) v')
X = torch.concat(data, dim=2)[:, SKIP::STRIDE, ::DECIMATE]
y = ds['binary_ans_base']-ds['binary_ans_adapt']
return dict(X=X, y=y)
def prepare_ds(ds):
"""
prepare a dataset for training
this should front load much of the computation
it should restrict it to the needed rows X and y
"""
ds = (ds
.with_format("torch")
.select_columns(input_columns)
.map(ds2xy_batched, batched=True, batch_size=128,
remove_columns=input_columns)
)
return ds
def load_file_to_dm(f):
ds = Dataset.from_file(str(f1_val), in_memory=True).with_format("torch")
ds = filter_ds_to_known(ds, verbose=True, true_col='truth')
ds = prepare_ds(ds)
# limit size
MAX_SAMPLES = min(len(ds), MAX_ROWS*2)
ds = ds.select(range(0, MAX_SAMPLES))
dm = ActivationDataModule(ds, f.stem, batch_size=batch_size, num_workers=4)
dm.setup()
return dm
In [ ]:
dm = load_file_to_dm(f1_val)
# dm_ood = load_file_to_dm(f1_ood)In [9]:
dm.ds.select(range(33))['X'].shapeOut [9]:
torch.Size([33, 7, 11520, 2])
In [17]:
# 4x faster
x = dm.ds.select(range(33)).with_format(None)['X']
x = torch.FloatTensor(x)
x.shape, xOut [17]:
(torch.Size([33, 7, 11520, 2]),
tensor([[[[-3.2861e-01, -4.0918e-01],
[ 1.4946e+00, 9.8340e-01],
[ 1.0400e-01, -4.7083e-01],
...,
[-6.0919e-01, -3.0310e-01],
[-3.8916e-01, -3.4668e-01],
[ 1.2476e-01, 2.8721e-01]],
[[-2.5488e-01, 6.2500e-02],
[ 1.0659e+00, 7.0117e-01],
[ 2.9956e-01, 4.8975e-01],
...,
[ 1.8530e-01, 2.1820e-01],
[ 8.8791e-02, 2.3016e-01],
[-7.7881e-02, 8.8043e-03]],
[[ 1.3745e+00, 1.2827e+00],
[ 8.5352e-01, 1.0352e+00],
[ 8.8672e-01, 9.1992e-01],
...,
[-2.9950e-01, -2.2668e-01],
[ 5.4620e-01, 4.8962e-01],
[-1.2872e-01, -4.1565e-02]],
...,
[[-6.0059e-02, 4.2236e-01],
[ 1.0012e+00, 5.1855e-01],
[-3.1559e+00, -2.0396e+00],
...,
[ 2.5513e-01, 1.5808e-01],
[ 3.6621e-02, 7.8094e-02],
[-6.4600e-01, -4.5972e-01]],
[[-3.6963e-01, -3.3447e-01],
[ 7.2510e-01, 7.0557e-01],
[ 3.2227e-02, -1.6504e-01],
...,
[ 1.1761e-01, 1.0216e-01],
[-2.6520e-01, -1.4828e-01],
[-6.3159e-01, -4.2773e-01]],
[[ 1.5781e+00, 1.7319e+00],
[-1.4648e-02, 9.5703e-02],
[-2.1094e-01, -2.1289e-01],
...,
[-1.8561e-01, -2.0541e-01],
[ 3.4290e-01, 4.2151e-01],
[-5.2924e-01, -5.6805e-01]]],
[[[-1.1963e-01, -2.5977e-01],
[ 5.2734e-01, 2.0215e-01],
[ 2.4719e-01, -2.7393e-01],
...,
[-5.2087e-01, -4.3396e-01],
[ 3.9062e-03, -7.4951e-02],
[ 8.0383e-02, 1.5494e-01]],
[[-5.4004e-01, -3.4473e-01],
[ 7.4463e-01, 9.1406e-01],
[ 1.2109e-01, 1.3574e-01],
...,
[-4.0698e-01, -9.0332e-02],
[-3.1250e-02, 1.7212e-01],
[-3.2959e-02, 9.5703e-02]],
[[ 1.1011e+00, 1.1162e+00],
[-7.8125e-03, 4.6582e-01],
[ 5.4102e-01, 1.2471e+00],
...,
[-7.1228e-02, -3.7781e-02],
[-4.9011e-02, 2.7527e-01],
[-3.6377e-02, -8.9050e-02]],
...,
[[-1.5537e+00, 3.7793e-01],
[-6.6992e-01, -9.1309e-02],
[-5.3418e-01, -9.3018e-01],
...,
[-1.3845e+00, -4.4055e-01],
[ 4.5557e-01, 4.3445e-01],
[-7.5195e-02, 4.0967e-01]],
[[ 3.1250e-02, -2.3828e-01],
[ 8.1543e-01, 6.0156e-01],
[ 1.7480e-01, -2.8418e-01],
...,
[-1.5826e-01, 2.9382e-01],
[-5.1599e-01, -1.0107e-01],
[-5.1025e-02, 2.8442e-02]],
[[ 2.4789e+00, 2.7275e+00],
[ 9.7168e-01, 1.0186e+00],
[ 2.1924e-01, 1.8506e-01],
...,
[-2.1332e-01, -2.7539e-01],
[-1.7960e-01, -1.8269e-01],
[-5.6931e-01, -5.4858e-01]]],
[[[ 9.5215e-02, -2.8564e-01],
[-1.6064e-01, -6.5625e-01],
[ 5.9595e-01, 3.2660e-01],
...,
[ 5.1715e-01, 4.2908e-01],
[-1.5558e-01, -3.6749e-01],
[ 4.2297e-01, 1.7407e-01]],
[[-5.9570e-01, -5.5469e-01],
[ 5.8691e-01, 2.1387e-01],
[ 1.5088e-01, 2.9785e-01],
...,
[ 1.6260e-01, 1.0156e-01],
[ 5.1318e-01, 6.9434e-01],
[ 2.9700e-01, 5.2625e-01]],
[[ 9.0234e-01, 1.0049e+00],
[ 1.8223e+00, 1.4531e+00],
[ 1.1670e+00, 1.3564e+00],
...,
[-3.5938e-01, -2.0306e-01],
[ 2.5037e-01, 3.9404e-01],
[ 1.6663e-02, -2.9724e-02]],
...,
[[ 5.9570e-02, 8.7256e-01],
[ 6.4502e-01, 2.8809e-01],
[-2.6475e+00, -1.9377e+00],
...,
[-8.3325e-01, -7.4988e-01],
[ 2.5940e-02, -2.3748e-01],
[-1.6235e-01, 3.0518e-03]],
[[ 4.9805e-02, -2.1436e-01],
[ 1.8201e+00, 1.4751e+00],
[ 1.4648e-01, 2.0410e-01],
...,
[ 6.7126e-01, 5.3821e-01],
[-4.5642e-01, -2.8931e-01],
[-8.4595e-01, -1.0127e+00]],
[[ 7.1680e-01, 1.0791e+00],
[ 4.2969e-01, 5.4785e-01],
[-1.0742e-01, -6.1035e-02],
...,
[-1.2976e-01, -1.7480e-01],
[-2.5861e-01, -1.3147e-01],
[-6.1316e-01, -5.8118e-01]]],
...,
[[[-1.1238e+00, -1.0891e+00],
[ 9.4507e-01, 4.9829e-01],
[ 7.9346e-01, 3.4974e-01],
...,
[-2.4584e-01, -1.4879e-01],
[-6.0504e-01, -6.3129e-01],
[-3.2506e-01, -9.6436e-02]],
[[-1.2793e-01, 2.5977e-01],
[ 7.5635e-01, 6.6260e-01],
[-7.6172e-02, 2.3145e-01],
...,
[ 5.4321e-02, -1.8201e-01],
[-3.3862e-01, -8.8074e-02],
[-1.0010e-02, -3.4912e-02]],
[[ 1.4189e+00, 1.4229e+00],
[ 1.2734e+00, 1.3398e+00],
[ 1.1260e+00, 1.4790e+00],
...,
[ 3.7048e-02, 1.0654e-01],
[ 1.1804e-01, 2.2211e-01],
[ 1.3242e-01, -1.3733e-02]],
...,
[[-8.5449e-02, -2.7295e-01],
[-4.5557e-01, -2.8662e-01],
[-1.2246e+00, -9.8535e-01],
...,
[-5.5362e-01, -4.7339e-01],
[ 2.7875e-01, 1.8005e-01],
[ 5.1086e-01, 2.1124e-01]],
[[ 4.6680e-01, 6.4746e-01],
[ 9.9658e-01, 1.0630e+00],
[-4.3457e-01, -8.4570e-01],
...,
[-1.2659e-01, 2.4962e-01],
[ 2.4835e-01, 2.0551e-01],
[-1.3721e-01, 2.8320e-02]],
[[ 2.6445e+00, 2.7156e+00],
[ 1.9028e+00, 1.7930e+00],
[-6.3721e-01, -6.1426e-01],
...,
[-1.2360e-01, -2.3798e-01],
[-3.0713e-01, -2.8662e-01],
[-3.9154e-01, -3.9203e-01]]],
[[[-7.8662e-01, -9.9487e-01],
[ 2.6221e-01, 8.5205e-02],
[ 1.4875e+00, 1.0596e+00],
...,
[-4.9280e-01, -3.0591e-01],
[-2.1948e-01, -4.2414e-01],
[ 2.3242e-01, 3.4802e-01]],
[[-2.7148e-01, -2.9004e-01],
[ 1.4082e+00, 1.0684e+00],
[ 5.9479e-01, 6.0449e-01],
...,
[-3.8788e-02, 2.9260e-01],
[-2.2357e-01, 9.4360e-02],
[ 2.7856e-01, 4.8547e-01]],
[[ 1.5137e-01, 6.2109e-01],
[ 1.0928e+00, 6.7578e-01],
[ 1.3711e+00, 1.5513e+00],
...,
[-3.5443e-01, -3.7183e-01],
[ 2.2736e-02, 2.5436e-01],
[-5.0592e-01, -6.2805e-02]],
...,
[[-2.8540e-01, 3.6182e-01],
[-2.6709e-01, 8.4961e-02],
[-1.4053e+00, -1.0566e+00],
...,
[ 1.3965e-01, -3.0273e-01],
[ 2.6245e-01, 6.3940e-01],
[ 2.7026e-01, 1.6650e-01]],
[[ 2.9297e-02, -1.9531e-03],
[ 8.7598e-01, 8.3008e-01],
[ 1.6455e+00, 1.1543e+00],
...,
[ 1.8993e-01, 8.5190e-02],
[-1.5338e-01, -3.4247e-01],
[ 1.5015e-02, 4.3457e-02]],
[[ 1.6958e+00, 2.0638e+00],
[ 1.4092e+00, 1.4062e+00],
[ 6.2427e-01, 5.6250e-01],
...,
[-1.6839e-01, -2.8098e-01],
[-5.3711e-02, 2.8809e-02],
[-7.6782e-01, -7.3810e-01]]],
[[[-1.1572e-01, -3.6426e-01],
[-1.6797e-01, -5.9131e-01],
[ 6.3397e-01, 2.5489e-01],
...,
[ 3.5620e-01, 1.8689e-01],
[-3.0688e-01, -3.6542e-01],
[ 2.8603e-01, 2.1411e-01]],
[[-1.3281e-01, 4.0039e-02],
[-2.5195e-01, -3.5156e-01],
[ 1.4062e-01, 5.1978e-01],
...,
[-4.0811e-01, -1.5015e-01],
[ 3.5254e-01, 7.4817e-01],
[ 5.1781e-02, 2.1603e-01]],
[[ 1.0830e+00, 9.4824e-01],
[ 1.0049e+00, 1.1729e+00],
[ 1.0576e+00, 1.1484e+00],
...,
[-5.1300e-01, -2.9749e-01],
[ 1.5605e-01, 3.9874e-01],
[-6.3477e-03, -1.0800e-01]],
...,
[[ 6.0645e-01, 1.2046e+00],
[-9.5703e-02, 1.9531e-03],
[-1.6113e+00, -1.7534e+00],
...,
[ 3.1628e-01, 2.9956e-01],
[ 2.8545e-01, 9.3872e-02],
[ 1.1011e-01, 2.9541e-02]],
[[ 4.5020e-01, -8.8867e-02],
[ 1.2534e+00, 9.7266e-01],
[ 5.0098e-01, 1.8457e-01],
...,
[ 7.2168e-01, -2.4463e-01],
[-5.8105e-01, -3.1342e-02],
[-8.3606e-01, -4.4604e-01]],
[[ 4.1699e-01, 8.1543e-01],
[ 8.5547e-01, 9.2188e-01],
[ 5.0098e-01, 5.6836e-01],
...,
[ 9.3994e-03, -7.7553e-02],
[-2.7484e-01, -9.8022e-02],
[-4.3964e-01, -4.1772e-01]]]]))[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.
In [12]:
dm.ds.select(range(33)).with_format('pt')['X'].shapeOut [12]:
torch.Size([33, 7, 11520, 2])
In [ ]: