Files
2023-10-15 20:17:36 +08:00

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)
2023-10-15 17:26:06.435 | INFO     | src.models.load:verbose_change_param:18 - changing pad_token_id from 32000 to 0
2023-10-15 17:26:06.435 | INFO     | src.models.load:verbose_change_param:18 - changing padding_side from right to left
2023-10-15 17:26:06.436 | INFO     | src.models.load:verbose_change_param:18 - changing truncation_side from right to left
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]
---------------------------------------------------------------------------
ArrowTypeError                            Traceback (most recent call last)
File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1703, in GeneratorBasedBuilder._prepare_split_single(self, gen_kwargs, fpath, file_format, max_shard_size, split_info, check_duplicate_keys, job_id)
   1702 num_shards = shard_id + 1
-> 1703 num_examples, num_bytes = writer.finalize()
   1704 writer.close()

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_writer.py:586, in ArrowWriter.finalize(self, close_stream)
    585     self.hkey_record = []
--> 586 self.write_examples_on_file()
    587 # If schema is known, infer features even if no examples were written

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_writer.py:448, in ArrowWriter.write_examples_on_file(self)
    444         batch_examples[col] = [
    445             row[0][col].to_pylist()[0] if isinstance(row[0][col], (pa.Array, pa.ChunkedArray)) else row[0][col]
    446             for row in self.current_examples
    447         ]
--> 448 self.write_batch(batch_examples=batch_examples)
    449 self.current_examples = []

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_writer.py:555, in ArrowWriter.write_batch(self, batch_examples, writer_batch_size)
    554 typed_sequence = OptimizedTypedSequence(col_values, type=col_type, try_type=col_try_type, col=col)
--> 555 arrays.append(pa.array(typed_sequence))
    556 inferred_features[col] = typed_sequence.get_inferred_type()

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/array.pxi:243, in pyarrow.lib.array()

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/array.pxi:110, in pyarrow.lib._handle_arrow_array_protocol()

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_writer.py:189, in TypedSequence.__arrow_array__(self, type)
    188     trying_cast_to_python_objects = True
--> 189     out = pa.array(cast_to_python_objects(data, only_1d_for_numpy=True))
    190 # use smaller integer precisions if possible

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/array.pxi:327, in pyarrow.lib.array()

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/array.pxi:39, in pyarrow.lib._sequence_to_array()

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/error.pxi:144, in pyarrow.lib.pyarrow_internal_check_status()

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/pyarrow/error.pxi:123, in pyarrow.lib.check_status()

ArrowTypeError: Expected bytes, got a 'list' object

The above exception was the direct cause of the following exception:

DatasetGenerationError                    Traceback (most recent call last)
/home/ubuntu/Documents/mjc/elk/discovering_latent_knowledge2/notebooks/012b_scratch_dataset.ipynb Cell 9 line 1
----> <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> ds1 = Dataset.from_generator(
      <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>     generator=batch_hidden_states,
      <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>     info=DatasetInfo(
      <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>         description=json.dumps(info_kwargs, indent=2),
      <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>         config_name=f,
      <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>     ),
      <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>     gen_kwargs=gen_kwargs,
      <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>     num_proc=1,
      <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> )
     <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> ds1

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/arrow_dataset.py:1072, in Dataset.from_generator(generator, features, cache_dir, keep_in_memory, gen_kwargs, num_proc, **kwargs)
   1016 """Create a Dataset from a generator.
   1017 
   1018 Args:
   (...)
   1060 ```
   1061 """
   1062 from .io.generator import GeneratorDatasetInputStream
   1064 return GeneratorDatasetInputStream(
   1065     generator=generator,
   1066     features=features,
   1067     cache_dir=cache_dir,
   1068     keep_in_memory=keep_in_memory,
   1069     gen_kwargs=gen_kwargs,
   1070     num_proc=num_proc,
   1071     **kwargs,
-> 1072 ).read()

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/io/generator.py:47, in GeneratorDatasetInputStream.read(self)
     44     verification_mode = None
     45     base_path = None
---> 47     self.builder.download_and_prepare(
     48         download_config=download_config,
     49         download_mode=download_mode,
     50         verification_mode=verification_mode,
     51         # try_from_hf_gcs=try_from_hf_gcs,
     52         base_path=base_path,
     53         num_proc=self.num_proc,
     54     )
     55     dataset = self.builder.as_dataset(
     56         split="train", verification_mode=verification_mode, in_memory=self.keep_in_memory
     57     )
     58 return dataset

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:954, in DatasetBuilder.download_and_prepare(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)
    952     if num_proc is not None:
    953         prepare_split_kwargs["num_proc"] = num_proc
--> 954     self._download_and_prepare(
    955         dl_manager=dl_manager,
    956         verification_mode=verification_mode,
    957         **prepare_split_kwargs,
    958         **download_and_prepare_kwargs,
    959     )
    960 # Sync info
    961 self.info.dataset_size = sum(split.num_bytes for split in self.info.splits.values())

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1717, in GeneratorBasedBuilder._download_and_prepare(self, dl_manager, verification_mode, **prepare_splits_kwargs)
   1716 def _download_and_prepare(self, dl_manager, verification_mode, **prepare_splits_kwargs):
-> 1717     super()._download_and_prepare(
   1718         dl_manager,
   1719         verification_mode,
   1720         check_duplicate_keys=verification_mode == VerificationMode.BASIC_CHECKS
   1721         or verification_mode == VerificationMode.ALL_CHECKS,
   1722         **prepare_splits_kwargs,
   1723     )

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1049, in DatasetBuilder._download_and_prepare(self, dl_manager, verification_mode, **prepare_split_kwargs)
   1045 split_dict.add(split_generator.split_info)
   1047 try:
   1048     # Prepare split will record examples associated to the split
-> 1049     self._prepare_split(split_generator, **prepare_split_kwargs)
   1050 except OSError as e:
   1051     raise OSError(
   1052         "Cannot find data file. "
   1053         + (self.manual_download_instructions or "")
   1054         + "\nOriginal error:\n"
   1055         + str(e)
   1056     ) from None

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1555, in GeneratorBasedBuilder._prepare_split(self, split_generator, check_duplicate_keys, file_format, num_proc, max_shard_size)
   1553 job_id = 0
   1554 with pbar:
-> 1555     for job_id, done, content in self._prepare_split_single(
   1556         gen_kwargs=gen_kwargs, job_id=job_id, **_prepare_split_args
   1557     ):
   1558         if done:
   1559             result = content

File ~/mambaforge/envs/dlk4/lib/python3.11/site-packages/datasets/builder.py:1712, in GeneratorBasedBuilder._prepare_split_single(self, gen_kwargs, fpath, file_format, max_shard_size, split_info, check_duplicate_keys, job_id)
   1710     if isinstance(e, SchemaInferenceError) and e.__context__ is not None:
   1711         e = e.__context__
-> 1712     raise DatasetGenerationError("An error occurred while generating the dataset") from e
   1714 yield job_id, True, (total_num_examples, total_num_bytes, writer._features, num_shards, shard_lengths)

DatasetGenerationError: 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']

Scratch

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 [ ]: