mirror of
https://github.com/wassname/discovering_latent_knowledge.git
synced 2026-09-12 12:13:04 +08:00
37 KiB
37 KiB
In [ ]:
In [1]:
# autoreload import your package
%load_ext autoreload
%autoreload 2
from make_dataset import *
In [2]:
cfg = ExtractConfig(max_examples=(10, 10), max_length=999)
cfg
Out [2]:
ExtractConfig(datasets=('amazon_polarity', 'super_glue:boolq', 'glue:qnli', 'imdb'), model='TheBloke/WizardCoder-Python-13B-V1.0-GPTQ', data_dirs=(), max_examples=(10, 10), num_shots=1, num_variants=-1, layers=(), seed=42, token_loc='last', template_path=None, max_length=999)In [3]:
model, tokenizer = load_model(cfg.model)
[32m2023-10-15 17:26:06.435[0m | [1mINFO [0m | [36msrc.models.load[0m:[36mverbose_change_param[0m:[36m18[0m - [1mchanging pad_token_id from 32000 to 0[0m [32m2023-10-15 17:26:06.435[0m | [1mINFO [0m | [36msrc.models.load[0m:[36mverbose_change_param[0m:[36m18[0m - [1mchanging padding_side from right to left[0m [32m2023-10-15 17:26:06.436[0m | [1mINFO [0m | [36msrc.models.load[0m:[36mverbose_change_param[0m:[36m18[0m - [1mchanging truncation_side from right to left[0m
In [4]:
ds_name = cfg.datasets[0]
split_type = "train"
ds_tokens = load_preproc_dataset(ds_name, cfg, tokenizer)
ds_tokens_calibration = ds_tokens.select(range(10))
Generating train split: 0 examples [00:00, ? examples/s]
Extracting 11 variants of each prompt
tokenize: 0%| | 0/32 [00:00<?, ? examples/s]
truncated: 0%| | 0/32 [00:00<?, ? examples/s]
prompt_truncated: 0%| | 0/32 [00:00<?, ? examples/s]
choice_ids: 0%| | 0/32 [00:00<?, ? examples/s]
Filter: 0%| | 0/32 [00:00<?, ? examples/s]
removed truncated rows to leave: num_rows 10
In [5]:
# b = next(iter(ds_tokens))
# b
In [6]:
# get the calibration dataset
BATCH_SIZE = 2
f = None
info_kwargs = dict(extract_cfg=cfg.to_dict(), ds_name=ds_name, split_type=split_type, f=f, date=pd.Timestamp.now().isoformat(),)
intervention_dicts = [None, ]
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,
)
In [ ]:
In [24]:
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,
)
ds1
Generating train split: 0 examples [00:00, ? examples/s]
get hidden states: 0%| | 0/5 [00:00<?, ?it/s]
[0;31m---------------------------------------------------------------------------[0m [0;31mArrowTypeError[0m Traceback (most recent call last) File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1703[0m, in [0;36mGeneratorBasedBuilder._prepare_split_single[0;34m(self, gen_kwargs, fpath, file_format, max_shard_size, split_info, check_duplicate_keys, job_id)[0m [1;32m 1702[0m num_shards [39m=[39m shard_id [39m+[39m [39m1[39m [0;32m-> 1703[0m num_examples, num_bytes [39m=[39m writer[39m.[39;49mfinalize() [1;32m 1704[0m writer[39m.[39mclose() File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_writer.py:586[0m, in [0;36mArrowWriter.finalize[0;34m(self, close_stream)[0m [1;32m 585[0m [39mself[39m[39m.[39mhkey_record [39m=[39m [] [0;32m--> 586[0m [39mself[39;49m[39m.[39;49mwrite_examples_on_file() [1;32m 587[0m [39m# If schema is known, infer features even if no examples were written[39;00m File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_writer.py:448[0m, in [0;36mArrowWriter.write_examples_on_file[0;34m(self)[0m [1;32m 444[0m batch_examples[col] [39m=[39m [ [1;32m 445[0m row[[39m0[39m][col][39m.[39mto_pylist()[[39m0[39m] [39mif[39;00m [39misinstance[39m(row[[39m0[39m][col], (pa[39m.[39mArray, pa[39m.[39mChunkedArray)) [39melse[39;00m row[[39m0[39m][col] [1;32m 446[0m [39mfor[39;00m row [39min[39;00m [39mself[39m[39m.[39mcurrent_examples [1;32m 447[0m ] [0;32m--> 448[0m [39mself[39;49m[39m.[39;49mwrite_batch(batch_examples[39m=[39;49mbatch_examples) [1;32m 449[0m [39mself[39m[39m.[39mcurrent_examples [39m=[39m [] File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_writer.py:555[0m, in [0;36mArrowWriter.write_batch[0;34m(self, batch_examples, writer_batch_size)[0m [1;32m 554[0m typed_sequence [39m=[39m OptimizedTypedSequence(col_values, [39mtype[39m[39m=[39mcol_type, try_type[39m=[39mcol_try_type, col[39m=[39mcol) [0;32m--> 555[0m arrays[39m.[39mappend(pa[39m.[39;49marray(typed_sequence)) [1;32m 556[0m inferred_features[col] [39m=[39m typed_sequence[39m.[39mget_inferred_type() File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/array.pxi:243[0m, in [0;36mpyarrow.lib.array[0;34m()[0m File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/array.pxi:110[0m, in [0;36mpyarrow.lib._handle_arrow_array_protocol[0;34m()[0m File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_writer.py:189[0m, in [0;36mTypedSequence.__arrow_array__[0;34m(self, type)[0m [1;32m 188[0m trying_cast_to_python_objects [39m=[39m [39mTrue[39;00m [0;32m--> 189[0m out [39m=[39m pa[39m.[39;49marray(cast_to_python_objects(data, only_1d_for_numpy[39m=[39;49m[39mTrue[39;49;00m)) [1;32m 190[0m [39m# use smaller integer precisions if possible[39;00m File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/array.pxi:327[0m, in [0;36mpyarrow.lib.array[0;34m()[0m File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/array.pxi:39[0m, in [0;36mpyarrow.lib._sequence_to_array[0;34m()[0m File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/error.pxi:144[0m, in [0;36mpyarrow.lib.pyarrow_internal_check_status[0;34m()[0m File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/error.pxi:123[0m, in [0;36mpyarrow.lib.check_status[0;34m()[0m [0;31mArrowTypeError[0m: Expected bytes, got a 'list' object The above exception was the direct cause of the following exception: [0;31mDatasetGenerationError[0m Traceback (most recent call last) [1;32m/home/ubuntu/Documents/mjc/elk/discovering_latent_knowledge2/notebooks/012b_scratch_dataset.ipynb Cell 9[0m line [0;36m1 [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>[0m ds1 [39m=[39m Dataset[39m.[39;49mfrom_generator( [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>[0m generator[39m=[39;49mbatch_hidden_states, [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>[0m info[39m=[39;49mDatasetInfo( [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>[0m description[39m=[39;49mjson[39m.[39;49mdumps(info_kwargs, indent[39m=[39;49m[39m2[39;49m), [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>[0m config_name[39m=[39;49mf, [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>[0m ), [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>[0m gen_kwargs[39m=[39;49mgen_kwargs, [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>[0m num_proc[39m=[39;49m[39m1[39;49m, [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>[0m ) [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>[0m ds1 File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_dataset.py:1072[0m, in [0;36mDataset.from_generator[0;34m(generator, features, cache_dir, keep_in_memory, gen_kwargs, num_proc, **kwargs)[0m [1;32m 1016[0m [39m[39m[39m"""Create a Dataset from a generator.[39;00m [1;32m 1017[0m [1;32m 1018[0m [39mArgs:[39;00m [0;32m (...)[0m [1;32m 1060[0m [39m```[39;00m [1;32m 1061[0m [39m"""[39;00m [1;32m 1062[0m [39mfrom[39;00m [39m.[39;00m[39mio[39;00m[39m.[39;00m[39mgenerator[39;00m [39mimport[39;00m GeneratorDatasetInputStream [1;32m 1064[0m [39mreturn[39;00m GeneratorDatasetInputStream( [1;32m 1065[0m generator[39m=[39;49mgenerator, [1;32m 1066[0m features[39m=[39;49mfeatures, [1;32m 1067[0m cache_dir[39m=[39;49mcache_dir, [1;32m 1068[0m keep_in_memory[39m=[39;49mkeep_in_memory, [1;32m 1069[0m gen_kwargs[39m=[39;49mgen_kwargs, [1;32m 1070[0m num_proc[39m=[39;49mnum_proc, [1;32m 1071[0m [39m*[39;49m[39m*[39;49mkwargs, [0;32m-> 1072[0m )[39m.[39;49mread() File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/io/generator.py:47[0m, in [0;36mGeneratorDatasetInputStream.read[0;34m(self)[0m [1;32m 44[0m verification_mode [39m=[39m [39mNone[39;00m [1;32m 45[0m base_path [39m=[39m [39mNone[39;00m [0;32m---> 47[0m [39mself[39;49m[39m.[39;49mbuilder[39m.[39;49mdownload_and_prepare( [1;32m 48[0m download_config[39m=[39;49mdownload_config, [1;32m 49[0m download_mode[39m=[39;49mdownload_mode, [1;32m 50[0m verification_mode[39m=[39;49mverification_mode, [1;32m 51[0m [39m# try_from_hf_gcs=try_from_hf_gcs,[39;49;00m [1;32m 52[0m base_path[39m=[39;49mbase_path, [1;32m 53[0m num_proc[39m=[39;49m[39mself[39;49m[39m.[39;49mnum_proc, [1;32m 54[0m ) [1;32m 55[0m dataset [39m=[39m [39mself[39m[39m.[39mbuilder[39m.[39mas_dataset( [1;32m 56[0m split[39m=[39m[39m"[39m[39mtrain[39m[39m"[39m, verification_mode[39m=[39mverification_mode, in_memory[39m=[39m[39mself[39m[39m.[39mkeep_in_memory [1;32m 57[0m ) [1;32m 58[0m [39mreturn[39;00m dataset File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:954[0m, in [0;36mDatasetBuilder.download_and_prepare[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)[0m [1;32m 952[0m [39mif[39;00m num_proc [39mis[39;00m [39mnot[39;00m [39mNone[39;00m: [1;32m 953[0m prepare_split_kwargs[[39m"[39m[39mnum_proc[39m[39m"[39m] [39m=[39m num_proc [0;32m--> 954[0m [39mself[39;49m[39m.[39;49m_download_and_prepare( [1;32m 955[0m dl_manager[39m=[39;49mdl_manager, [1;32m 956[0m verification_mode[39m=[39;49mverification_mode, [1;32m 957[0m [39m*[39;49m[39m*[39;49mprepare_split_kwargs, [1;32m 958[0m [39m*[39;49m[39m*[39;49mdownload_and_prepare_kwargs, [1;32m 959[0m ) [1;32m 960[0m [39m# Sync info[39;00m [1;32m 961[0m [39mself[39m[39m.[39minfo[39m.[39mdataset_size [39m=[39m [39msum[39m(split[39m.[39mnum_bytes [39mfor[39;00m split [39min[39;00m [39mself[39m[39m.[39minfo[39m.[39msplits[39m.[39mvalues()) File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1717[0m, in [0;36mGeneratorBasedBuilder._download_and_prepare[0;34m(self, dl_manager, verification_mode, **prepare_splits_kwargs)[0m [1;32m 1716[0m [39mdef[39;00m [39m_download_and_prepare[39m([39mself[39m, dl_manager, verification_mode, [39m*[39m[39m*[39mprepare_splits_kwargs): [0;32m-> 1717[0m [39msuper[39;49m()[39m.[39;49m_download_and_prepare( [1;32m 1718[0m dl_manager, [1;32m 1719[0m verification_mode, [1;32m 1720[0m check_duplicate_keys[39m=[39;49mverification_mode [39m==[39;49m VerificationMode[39m.[39;49mBASIC_CHECKS [1;32m 1721[0m [39mor[39;49;00m verification_mode [39m==[39;49m VerificationMode[39m.[39;49mALL_CHECKS, [1;32m 1722[0m [39m*[39;49m[39m*[39;49mprepare_splits_kwargs, [1;32m 1723[0m ) File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1049[0m, in [0;36mDatasetBuilder._download_and_prepare[0;34m(self, dl_manager, verification_mode, **prepare_split_kwargs)[0m [1;32m 1045[0m split_dict[39m.[39madd(split_generator[39m.[39msplit_info) [1;32m 1047[0m [39mtry[39;00m: [1;32m 1048[0m [39m# Prepare split will record examples associated to the split[39;00m [0;32m-> 1049[0m [39mself[39;49m[39m.[39;49m_prepare_split(split_generator, [39m*[39;49m[39m*[39;49mprepare_split_kwargs) [1;32m 1050[0m [39mexcept[39;00m [39mOSError[39;00m [39mas[39;00m e: [1;32m 1051[0m [39mraise[39;00m [39mOSError[39;00m( [1;32m 1052[0m [39m"[39m[39mCannot find data file. [39m[39m"[39m [1;32m 1053[0m [39m+[39m ([39mself[39m[39m.[39mmanual_download_instructions [39mor[39;00m [39m"[39m[39m"[39m) [1;32m 1054[0m [39m+[39m [39m"[39m[39m\n[39;00m[39mOriginal error:[39m[39m\n[39;00m[39m"[39m [1;32m 1055[0m [39m+[39m [39mstr[39m(e) [1;32m 1056[0m ) [39mfrom[39;00m [39mNone[39;00m File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1555[0m, in [0;36mGeneratorBasedBuilder._prepare_split[0;34m(self, split_generator, check_duplicate_keys, file_format, num_proc, max_shard_size)[0m [1;32m 1553[0m job_id [39m=[39m [39m0[39m [1;32m 1554[0m [39mwith[39;00m pbar: [0;32m-> 1555[0m [39mfor[39;00m job_id, done, content [39min[39;00m [39mself[39m[39m.[39m_prepare_split_single( [1;32m 1556[0m gen_kwargs[39m=[39mgen_kwargs, job_id[39m=[39mjob_id, [39m*[39m[39m*[39m_prepare_split_args [1;32m 1557[0m ): [1;32m 1558[0m [39mif[39;00m done: [1;32m 1559[0m result [39m=[39m content File [0;32m~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1712[0m, in [0;36mGeneratorBasedBuilder._prepare_split_single[0;34m(self, gen_kwargs, fpath, file_format, max_shard_size, split_info, check_duplicate_keys, job_id)[0m [1;32m 1710[0m [39mif[39;00m [39misinstance[39m(e, SchemaInferenceError) [39mand[39;00m e[39m.[39m__context__ [39mis[39;00m [39mnot[39;00m [39mNone[39;00m: [1;32m 1711[0m e [39m=[39m e[39m.[39m__context__ [0;32m-> 1712[0m [39mraise[39;00m DatasetGenerationError([39m"[39m[39mAn error occurred while generating the dataset[39m[39m"[39m) [39mfrom[39;00m [39me[39;00m [1;32m 1714[0m [39myield[39;00m job_id, [39mTrue[39;00m, (total_num_examples, total_num_bytes, writer[39m.[39m_features, num_shards, shard_lengths) [0;31mDatasetGenerationError[0m: An error occurred while generating the dataset
In [ ]:
ds1.info.description
ds1['layer_names']
['model.layers.4.self_attn', 'model.layers.8.self_attn', 'model.layers.4.self_attn', 'model.layers.8.self_attn', 'model.layers.4.self_attn', 'model.layers.8.self_attn', 'model.layers.4.self_attn', 'model.layers.8.self_attn', 'model.layers.4.self_attn', 'model.layers.8.self_attn']
In [8]:
from src.datasets.hs import ExtractHiddenStates
ehs = ExtractHiddenStates(model, tokenizer, intervention_dicts=intervention_dicts, layer_stride=cfg.layer_stride, layer_padding=cfg.layer_padding)
ehs
Out [8]:
ExtractHiddenStates(model=LlamaForCausalLM(
(model): LlamaModel(
(embed_tokens): Embedding(32001, 5120, padding_idx=0)
(layers): ModuleList(
(0-39): 40 x LlamaDecoderLayer(
(self_attn): LlamaAttention(
(rotary_emb): LlamaRotaryEmbedding()
(k_proj): QuantLinear()
(o_proj): QuantLinear()
(q_proj): QuantLinear()
(v_proj): QuantLinear()
)
(mlp): LlamaMLP(
(act_fn): SiLUActivation()
(down_proj): QuantLinear()
(gate_proj): QuantLinear()
(up_proj): QuantLinear()
)
(input_layernorm): LlamaRMSNorm()
(post_attention_layernorm): LlamaRMSNorm()
)
)
(norm): LlamaRMSNorm()
)
(lm_head): Linear(in_features=5120, out_features=32001, bias=False)
), tokenizer=LlamaTokenizerFast(name_or_path='TheBloke/WizardCoder-Python-13B-V1.0-GPTQ', vocab_size=32000, model_max_length=1000000000000000019884624838656, is_fast=True, padding_side='left', truncation_side='left', special_tokens={'bos_token': '</s>', 'eos_token': '</s>', 'unk_token': '</s>', 'pad_token': '<unk>'}, clean_up_tokenization_spaces=False), intervention_dicts=[None], layer_stride=4, layer_padding=4)In [9]:
from torch.utils.data import DataLoader
batch_size = BATCH_SIZE
data = ds_tokens_calibration
from src.helpers.ds import ds_keep_cols, clear_mem
# get a batch
torch_cols = ['input_ids', 'attention_mask', 'choice_ids']
ds_t_subset = ds_keep_cols(data, torch_cols)
ds_t_subset.set_format(type='torch')
ds_p_subset = data.remove_columns(torch_cols)
dl = DataLoader(ds_t_subset, batch_size=batch_size, shuffle=False)
for i, batch in enumerate(tqdm(dl, desc='get hidden states')):
input_ids, attention_mask, choice_ids = batch["input_ids"], batch["attention_mask"], batch["choice_ids"]
# nn = len(input_ids)
# index = i*batch_size+np.arange(nn)
# # different due to dropout
# hsl = ehs.get_batch_of_hidden_states(input_ids=input_ids, attention_mask=attention_mask, choice_ids=choice_ids)
get hidden states: 0%| | 0/5 [00:00<?, ?it/s]
In [10]:
# from src.datasets.intervene import get_com_directions, get_interventions_dict
# head_wise_activations = np.array(ds1['head_activation'])
# labels=np.array(ds1["label_true"])
# get_com_directions(2,
# 2,
# head_wise_activations,
# labels
# )
In [11]:
from einops import rearrange, reduce, repeat, asnumpy, parse_shape
from src.datasets.intervene import InterventionDict
from typing import Tuple
from functools import partial
from baukit.nethook import Trace, TraceDict, recursive_copy
from src.datasets.intervene import intervention_meta_fn, get_interventions_dict
activations = np.array(ds1['head_activation']).squeeze(-1)
labels = np.array(ds1["label_true"]).astype(int)==1
num_heads = model.config.num_attention_heads
layer_names = [f"model.layers.{i}.self_attn" for i in range(model.config.num_hidden_layers)]
layer_names, layer_inds = ehs.get_layer_selection(layer_names)
In [12]:
ds1.info.description
Out [12]:
'{\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}'In [13]:
interventions = get_interventions_dict(activations, labels, layer_names, num_heads)
intervention_fn = partial(intervention_meta_fn, interventions=interventions, num_heads=num_heads)
model.cuda().eval()
device = model.device
with torch.no_grad():
with TraceDict(model, layer_names, edit_output=intervention_fn) as ret:
outputs = model(input_ids=input_ids.to(device), attention_mask=attention_mask.to(device), return_dict=True, output_hidden_states=True)
a = outputs[0]
a
Out [13]:
tensor([[[-3.4297, 12.6875, 0.4692, ..., -2.2012, -1.2109, -2.1602],
[-3.4297, 12.6875, 0.4695, ..., -2.2012, -1.2129, -2.1621],
[-3.4336, 12.6875, 0.4702, ..., -2.2051, -1.2129, -2.1621],
...,
[-2.4434, 1.2295, 11.8359, ..., -2.0957, 0.1870, 0.2942],
[-4.2383, -4.3594, 12.7656, ..., -2.8594, -1.1748, -0.6260],
[-6.4336, -6.3516, 11.0547, ..., -5.6250, -2.6367, -2.8652]],
[[-3.4062, 12.5859, 0.4016, ..., -2.1797, -1.2871, -2.1543],
[-3.4082, 12.5625, 0.3992, ..., -2.1777, -1.2900, -2.1543],
[-3.4062, 12.5703, 0.3999, ..., -2.1777, -1.2910, -2.1543],
...,
[-2.4629, 0.2039, 12.3828, ..., -2.2266, 0.8071, 0.2996],
[-6.0508, -6.7305, 10.3047, ..., -4.7031, -2.3340, -2.0586],
[-7.2188, -8.3750, 9.3047, ..., -5.4141, -2.8828, -3.3223]]],
device='cuda:0')In [14]:
with torch.no_grad():
with TraceDict(model, layer_names) as ret:
outputs = model(input_ids=input_ids.to(device), attention_mask=attention_mask.to(device), return_dict=True, output_hidden_states=True)
a = outputs[0]
a
Out [14]:
tensor([[[-3.4277, 12.6953, 0.4692, ..., -2.2012, -1.2100, -2.1602],
[-3.4297, 12.6875, 0.4695, ..., -2.2012, -1.2119, -2.1602],
[-3.4297, 12.6953, 0.4705, ..., -2.2031, -1.2119, -2.1602],
...,
[-2.4434, 1.2295, 11.8359, ..., -2.0957, 0.1870, 0.2942],
[-4.2383, -4.3594, 12.7656, ..., -2.8594, -1.1748, -0.6260],
[-5.4531, -4.9375, 10.1641, ..., -4.8320, -2.1133, -2.2773]],
[[-3.4043, 12.5703, 0.4009, ..., -2.1777, -1.2881, -2.1523],
[-3.4043, 12.5781, 0.4006, ..., -2.1758, -1.2881, -2.1523],
[-3.4043, 12.5703, 0.4006, ..., -2.1777, -1.2900, -2.1543],
...,
[-2.4629, 0.2039, 12.3828, ..., -2.2266, 0.8071, 0.2996],
[-6.0508, -6.7305, 10.3047, ..., -4.7031, -2.3340, -2.0586],
[-6.3906, -7.0547, 8.7266, ..., -4.8984, -2.3691, -2.8887]]],
device='cuda:0')In [15]:
intervention_fn2 = partial(intervention_meta_fn, interventions=interventions, num_heads=num_heads, alpha=-15)
model.cuda().eval()
device = model.device
with torch.no_grad():
with TraceDict(model, layer_names, edit_output=intervention_fn2) as ret:
outputs = model(input_ids=input_ids.to(device), attention_mask=attention_mask.to(device), return_dict=True, output_hidden_states=True)
a = outputs[0]
a
Out [15]:
tensor([[[-3.4277, 12.6953, 0.4695, ..., -2.2012, -1.2090, -2.1602],
[-3.4277, 12.6875, 0.4697, ..., -2.1992, -1.2119, -2.1602],
[-3.4316, 12.6875, 0.4705, ..., -2.2051, -1.2119, -2.1621],
...,
[-2.4434, 1.2295, 11.8359, ..., -2.0957, 0.1870, 0.2942],
[-4.2383, -4.3594, 12.7656, ..., -2.8594, -1.1748, -0.6260],
[-4.9805, -4.4336, 8.6094, ..., -4.4102, -1.9639, -2.0254]],
[[-3.4062, 12.5859, 0.4014, ..., -2.1797, -1.2881, -2.1543],
[-3.4082, 12.5859, 0.4011, ..., -2.1797, -1.2881, -2.1562],
[-3.4062, 12.5703, 0.3999, ..., -2.1777, -1.2910, -2.1543],
...,
[-2.4629, 0.2039, 12.3828, ..., -2.2266, 0.8071, 0.2996],
[-6.0508, -6.7305, 10.3047, ..., -4.7031, -2.3340, -2.0586],
[-5.8555, -6.3125, 7.6172, ..., -4.5625, -2.0273, -2.6055]]],
device='cuda:0')In [ ]: