Files
2024-01-14 10:33:31 +08:00

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')

Paramsnet

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')

Load data

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_ood
In [ ]:
# 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'].shape
Out [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, x
Out [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]]]]))
The 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'].shape
Out [12]:
torch.Size([33, 7, 11520, 2])
In [ ]: