\u001b[0m:\u001b[94m33\u001b[0m \u001b[31m│\u001b[0m\n",
+ "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
+ "\u001b[31m│\u001b[0m \u001b[2m30 \u001b[0m\u001b[2m│ │ \u001b[0m]) \u001b[31m│\u001b[0m\n",
+ "\u001b[31m│\u001b[0m \u001b[2m31 \u001b[0m\u001b[2m│ \u001b[0m \u001b[31m│\u001b[0m\n",
+ "\u001b[31m│\u001b[0m \u001b[2m32 \u001b[0m\u001b[2m│ \u001b[0m\u001b[2m# FIXME not all the hidden state are the same size, wat\u001b[0m \u001b[31m│\u001b[0m\n",
+ "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m33 \u001b[2m│ \u001b[0mres = [np.concatenate(r) \u001b[94mfor\u001b[0m r \u001b[95min\u001b[0m \u001b[96mzip\u001b[0m(*res)] \u001b[31m│\u001b[0m\n",
+ "\u001b[31m│\u001b[0m \u001b[2m34 \u001b[0m\u001b[2m│ \u001b[0m\u001b[94mreturn\u001b[0m res \u001b[31m│\u001b[0m\n",
+ "\u001b[31m│\u001b[0m \u001b[2m35 \u001b[0m\u001b[2m│ \u001b[0mall_neg_hs, all_pos_hs, all_gt_labels, all_neg_ans, all_pos_ans = res \u001b[31m│\u001b[0m\n",
+ "\u001b[31m│\u001b[0m \u001b[2m36 \u001b[0m\u001b[2m│ \u001b[0m\u001b[94mreturn\u001b[0m all_neg_hs, all_pos_hs, all_gt_labels, all_neg_ans, all_pos_ans \u001b[31m│\u001b[0m\n",
+ "\u001b[31m│\u001b[0m in \u001b[92mconcatenate\u001b[0m:\u001b[94m200\u001b[0m \u001b[31m│\u001b[0m\n",
+ "\u001b[31m╰──────────────────────────────────────────────────────────────────────────────────────────────────╯\u001b[0m\n",
+ "\u001b[1;91mValueError: \u001b[0mall the input array dimensions except for the concatenation axis must match exactly, but along \n",
+ "dimension \u001b[1;36m1\u001b[0m, the array at index \u001b[1;36m0\u001b[0m has size \u001b[1;36m2285568\u001b[0m and the array at index \u001b[1;36m1\u001b[0m has size \u001b[1;36m1908736\u001b[0m\n"
]
},
- "execution_count": 13,
"metadata": {},
- "output_type": "execute_result"
+ "output_type": "display_data"
}
],
"source": [
"neg_hs, pos_hs, y, all_neg_ans, all_pos_ans = get_hidden_states_many_examples(model, tokenizer, data, model_type)\n",
"\n",
- "import gc\n",
+ "\n",
"gc.collect()\n",
"torch.cuda.empty_cache()\n",
"gc.collect()"
@@ -737,37 +1107,27 @@
"execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-14T06:40:56.451564Z",
- "start_time": "2023-05-14T06:40:56.451556Z"
+ "end_time": "2023-05-20T01:57:03.641742Z",
+ "start_time": "2023-05-20T01:57:03.641735Z"
}
},
"outputs": [],
- "source": []
+ "source": [
+ "# all_pos_ans"
+ ]
},
{
"cell_type": "code",
- "execution_count": 14,
+ "execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T00:29:24.238307Z",
- "start_time": "2023-05-19T00:29:24.223444Z"
+ "end_time": "2023-05-20T01:57:03.642523Z",
+ "start_time": "2023-05-20T01:57:03.642517Z"
}
},
- "outputs": [
- {
- "data": {
- "text/plain": [
- "(0.5366586538461539, 0.5759214743589743)"
- ]
- },
- "execution_count": 14,
- "metadata": {},
- "output_type": "execute_result"
- }
- ],
+ "outputs": [],
"source": [
- "from sklearn.metrics import f1_score, roc_auc_score, accuracy_score\n",
- "\n",
+ "# roc_auc_score\n",
"pos_score = roc_auc_score(y, all_pos_ans)\n",
"neg_score = roc_auc_score(y, all_neg_ans)\n",
"pos_score, neg_score"
@@ -787,28 +1147,17 @@
},
{
"cell_type": "code",
- "execution_count": 15,
+ "execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T00:29:24.250717Z",
- "start_time": "2023-05-19T00:29:24.239259Z"
+ "end_time": "2023-05-20T01:57:03.643517Z",
+ "start_time": "2023-05-20T01:57:03.643507Z"
},
"scrolled": true
},
- "outputs": [
- {
- "data": {
- "text/plain": [
- "(0.48, 0.52)"
- ]
- },
- "execution_count": 15,
- "metadata": {},
- "output_type": "execute_result"
- }
- ],
+ "outputs": [],
"source": [
- "\n",
+ "# accuracy_score\n",
"pos_score = accuracy_score(y, (all_pos_ans>0.)*1.0)\n",
"neg_score = accuracy_score(y, (all_neg_ans<0.5)*1.0)\n",
"pos_score, neg_score"
@@ -827,23 +1176,14 @@
},
{
"cell_type": "code",
- "execution_count": 16,
+ "execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T00:29:24.317414Z",
- "start_time": "2023-05-19T00:29:24.251856Z"
+ "end_time": "2023-05-20T01:57:03.644197Z",
+ "start_time": "2023-05-20T01:57:03.644190Z"
}
},
- "outputs": [
- {
- "name": "stdout",
- "output_type": "stream",
- "text": [
- "Logistic regression accuracy: 1.0 [TRAIN]\n",
- "Logistic regression accuracy: 0.96 [TEST]\n"
- ]
- }
- ],
+ "outputs": [],
"source": [
"# let's create a simple 50/50 train split (the data is already randomized)\n",
"n = len(y)\n",
@@ -887,11 +1227,11 @@
},
{
"cell_type": "code",
- "execution_count": 17,
+ "execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T00:29:24.321545Z",
- "start_time": "2023-05-19T00:29:24.318715Z"
+ "end_time": "2023-05-20T01:57:03.644851Z",
+ "start_time": "2023-05-20T01:57:03.644841Z"
}
},
"outputs": [],
@@ -942,11 +1282,11 @@
},
{
"cell_type": "code",
- "execution_count": 18,
+ "execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T00:29:24.336538Z",
- "start_time": "2023-05-19T00:29:24.322608Z"
+ "end_time": "2023-05-20T01:57:03.645452Z",
+ "start_time": "2023-05-20T01:57:03.645446Z"
}
},
"outputs": [],
@@ -965,11 +1305,11 @@
},
{
"cell_type": "code",
- "execution_count": 19,
+ "execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T00:29:24.357517Z",
- "start_time": "2023-05-19T00:29:24.337335Z"
+ "end_time": "2023-05-20T01:57:03.645991Z",
+ "start_time": "2023-05-20T01:57:03.645985Z"
}
},
"outputs": [],
@@ -1007,17 +1347,15 @@
},
{
"cell_type": "code",
- "execution_count": 20,
+ "execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T00:29:25.004935Z",
- "start_time": "2023-05-19T00:29:24.358571Z"
+ "end_time": "2023-05-19T04:12:55.004017Z",
+ "start_time": "2023-05-19T04:12:55.004011Z"
}
},
"outputs": [],
- "source": [
- "import lightning.pytorch as pl"
- ]
+ "source": []
},
{
"cell_type": "markdown",
@@ -1040,556 +1378,16 @@
},
{
"cell_type": "code",
- "execution_count": 21,
+ "execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T01:18:07.507525Z",
- "start_time": "2023-05-19T00:29:25.006004Z"
+ "end_time": "2023-05-20T01:57:03.646546Z",
+ "start_time": "2023-05-20T01:57:03.646539Z"
},
"scrolled": true
},
- "outputs": [
- {
- "name": "stderr",
- "output_type": "stream",
- "text": [
- "Found cached dataset amazon_polarity (/home/wassname/.cache/huggingface/datasets/amazon_polarity/amazon_polarity/3.0.0/a27b32b7e7b88eb274a8fa8ba0f654f1fe998a87c22547557317793b5d2772dc)\n",
- "get_hidden_states: 90%|███████████████████████████████████▊ | 895/1000 [08:21<00:58, 1.78examples/s]\n"
- ]
- },
- {
- "data": {
- "text/html": [
- "╭─────────────────────────────── Traceback (most recent call last) ────────────────────────────────╮\n",
- "│ in <cell line: 94>:94 │\n",
- "│ │\n",
- "│ 91 │\n",
- "│ 92 # test │\n",
- "│ 93 dm = IMBDHSDataModule(model, tokenizer) │\n",
- "│ ❱ 94 dm.setup('train') │\n",
- "│ 95 dl = dm.val_dataloader() │\n",
- "│ 96 b = next(iter(dl)) │\n",
- "│ 97 b │\n",
- "│ │\n",
- "│ in setup:39 │\n",
- "│ │\n",
- "│ 36 │ │ │\n",
- "│ 37 │ │ self.dataset = load_dataset(self.hparams.dataset_name, split=\"test\") │\n",
- "│ 38 │ │ │\n",
- "│ ❱ 39 │ │ neg_hs, pos_hs, y, all_neg_ans, all_pos_ans = get_hidden_states_many_examples( │\n",
- "│ 40 │ │ │ self.model, self.tokenizer, self.dataset, self.hparams.model_type, n=self.hp │\n",
- "│ 41 │ │ │\n",
- "│ 42 │ │ # let's create a simple 50/50 train split (the data is already randomized) │\n",
- "│ │\n",
- "│ in get_hidden_states_many_examples:27 │\n",
- "│ │\n",
- "│ 24 │ │ │\n",
- "│ 25 │ │ # get hidden states │\n",
- "│ 26 # print(format_imdb(text, 0)) │\n",
- "│ ❱ 27 │ │ neg = get_hidden_states(model, tokenizer, format_imdb(text, 0), model_type=model │\n",
- "│ 28 │ │ pos = get_hidden_states(model, tokenizer, format_imdb(text, 1), model_type=model │\n",
- "│ 29 │ │ │\n",
- "│ 30 │ │ # collect │\n",
- "│ │\n",
- "│ in get_hidden_states:96 │\n",
- "│ │\n",
- "│ 93 # \"encoder\": get_encoder_hidden_states, \"encoder_decoder\": get_encoder_decoder_hidden_st │\n",
- "│ 94 │ │ \"decoder\": get_decoder_hidden_states}[model_type] │\n",
- "│ 95 │ │\n",
- "│ ❱ 96 │ return fn(model, tokenizer, input_text, layers=layers) │\n",
- "│ 97 │\n",
- "│ │\n",
- "│ in get_decoder_hidden_states:56 │\n",
- "│ │\n",
- "│ 53 │ │\n",
- "│ 54 │ with torch.no_grad(): │\n",
- "│ 55 │ │ # FIXME: should be a batch, to speed it up │\n",
- "│ ❱ 56 │ │ output = model(input_ids, │\n",
- "│ 57 │ │ │ │ │ output_hidden_states=True │\n",
- "│ 58 # , output_attentions=True │\n",
- "│ 59 │ │ │ │ │ ) │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/torch/nn/modules/module.py:1501 │\n",
- "│ in _call_impl │\n",
- "│ │\n",
- "│ 1498 │ │ if not (self._backward_hooks or self._backward_pre_hooks or self._forward_hooks │\n",
- "│ 1499 │ │ │ │ or _global_backward_pre_hooks or _global_backward_hooks │\n",
- "│ 1500 │ │ │ │ or _global_forward_hooks or _global_forward_pre_hooks): │\n",
- "│ ❱ 1501 │ │ │ return forward_call(*args, **kwargs) │\n",
- "│ 1502 │ │ # Do not call functions when jit is used │\n",
- "│ 1503 │ │ full_backward_hooks, non_full_backward_hooks = [], [] │\n",
- "│ 1504 │ │ backward_pre_hooks = [] │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/accelerate/hooks.py:165 in │\n",
- "│ new_forward │\n",
- "│ │\n",
- "│ 162 │ │ │ with torch.no_grad(): │\n",
- "│ 163 │ │ │ │ output = old_forward(*args, **kwargs) │\n",
- "│ 164 │ │ else: │\n",
- "│ ❱ 165 │ │ │ output = old_forward(*args, **kwargs) │\n",
- "│ 166 │ │ return module._hf_hook.post_forward(module, output) │\n",
- "│ 167 │ │\n",
- "│ 168 │ module.forward = new_forward │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/transformers/models/llama/modeli │\n",
- "│ ng_llama.py:687 in forward │\n",
- "│ │\n",
- "│ 684 │ │ return_dict = return_dict if return_dict is not None else self.config.use_return │\n",
- "│ 685 │ │ │\n",
- "│ 686 │ │ # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn) │\n",
- "│ ❱ 687 │ │ outputs = self.model( │\n",
- "│ 688 │ │ │ input_ids=input_ids, │\n",
- "│ 689 │ │ │ attention_mask=attention_mask, │\n",
- "│ 690 │ │ │ position_ids=position_ids, │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/torch/nn/modules/module.py:1501 │\n",
- "│ in _call_impl │\n",
- "│ │\n",
- "│ 1498 │ │ if not (self._backward_hooks or self._backward_pre_hooks or self._forward_hooks │\n",
- "│ 1499 │ │ │ │ or _global_backward_pre_hooks or _global_backward_hooks │\n",
- "│ 1500 │ │ │ │ or _global_forward_hooks or _global_forward_pre_hooks): │\n",
- "│ ❱ 1501 │ │ │ return forward_call(*args, **kwargs) │\n",
- "│ 1502 │ │ # Do not call functions when jit is used │\n",
- "│ 1503 │ │ full_backward_hooks, non_full_backward_hooks = [], [] │\n",
- "│ 1504 │ │ backward_pre_hooks = [] │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/accelerate/hooks.py:165 in │\n",
- "│ new_forward │\n",
- "│ │\n",
- "│ 162 │ │ │ with torch.no_grad(): │\n",
- "│ 163 │ │ │ │ output = old_forward(*args, **kwargs) │\n",
- "│ 164 │ │ else: │\n",
- "│ ❱ 165 │ │ │ output = old_forward(*args, **kwargs) │\n",
- "│ 166 │ │ return module._hf_hook.post_forward(module, output) │\n",
- "│ 167 │ │\n",
- "│ 168 │ module.forward = new_forward │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/transformers/models/llama/modeli │\n",
- "│ ng_llama.py:577 in forward │\n",
- "│ │\n",
- "│ 574 │ │ │ │ │ None, │\n",
- "│ 575 │ │ │ │ ) │\n",
- "│ 576 │ │ │ else: │\n",
- "│ ❱ 577 │ │ │ │ layer_outputs = decoder_layer( │\n",
- "│ 578 │ │ │ │ │ hidden_states, │\n",
- "│ 579 │ │ │ │ │ attention_mask=attention_mask, │\n",
- "│ 580 │ │ │ │ │ position_ids=position_ids, │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/torch/nn/modules/module.py:1501 │\n",
- "│ in _call_impl │\n",
- "│ │\n",
- "│ 1498 │ │ if not (self._backward_hooks or self._backward_pre_hooks or self._forward_hooks │\n",
- "│ 1499 │ │ │ │ or _global_backward_pre_hooks or _global_backward_hooks │\n",
- "│ 1500 │ │ │ │ or _global_forward_hooks or _global_forward_pre_hooks): │\n",
- "│ ❱ 1501 │ │ │ return forward_call(*args, **kwargs) │\n",
- "│ 1502 │ │ # Do not call functions when jit is used │\n",
- "│ 1503 │ │ full_backward_hooks, non_full_backward_hooks = [], [] │\n",
- "│ 1504 │ │ backward_pre_hooks = [] │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/accelerate/hooks.py:165 in │\n",
- "│ new_forward │\n",
- "│ │\n",
- "│ 162 │ │ │ with torch.no_grad(): │\n",
- "│ 163 │ │ │ │ output = old_forward(*args, **kwargs) │\n",
- "│ 164 │ │ else: │\n",
- "│ ❱ 165 │ │ │ output = old_forward(*args, **kwargs) │\n",
- "│ 166 │ │ return module._hf_hook.post_forward(module, output) │\n",
- "│ 167 │ │\n",
- "│ 168 │ module.forward = new_forward │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/transformers/models/llama/modeli │\n",
- "│ ng_llama.py:305 in forward │\n",
- "│ │\n",
- "│ 302 │ │ # Fully Connected │\n",
- "│ 303 │ │ residual = hidden_states │\n",
- "│ 304 │ │ hidden_states = self.post_attention_layernorm(hidden_states) │\n",
- "│ ❱ 305 │ │ hidden_states = self.mlp(hidden_states) │\n",
- "│ 306 │ │ hidden_states = residual + hidden_states │\n",
- "│ 307 │ │ │\n",
- "│ 308 │ │ outputs = (hidden_states,) │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/torch/nn/modules/module.py:1501 │\n",
- "│ in _call_impl │\n",
- "│ │\n",
- "│ 1498 │ │ if not (self._backward_hooks or self._backward_pre_hooks or self._forward_hooks │\n",
- "│ 1499 │ │ │ │ or _global_backward_pre_hooks or _global_backward_hooks │\n",
- "│ 1500 │ │ │ │ or _global_forward_hooks or _global_forward_pre_hooks): │\n",
- "│ ❱ 1501 │ │ │ return forward_call(*args, **kwargs) │\n",
- "│ 1502 │ │ # Do not call functions when jit is used │\n",
- "│ 1503 │ │ full_backward_hooks, non_full_backward_hooks = [], [] │\n",
- "│ 1504 │ │ backward_pre_hooks = [] │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/accelerate/hooks.py:165 in │\n",
- "│ new_forward │\n",
- "│ │\n",
- "│ 162 │ │ │ with torch.no_grad(): │\n",
- "│ 163 │ │ │ │ output = old_forward(*args, **kwargs) │\n",
- "│ 164 │ │ else: │\n",
- "│ ❱ 165 │ │ │ output = old_forward(*args, **kwargs) │\n",
- "│ 166 │ │ return module._hf_hook.post_forward(module, output) │\n",
- "│ 167 │ │\n",
- "│ 168 │ module.forward = new_forward │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/transformers/models/llama/modeli │\n",
- "│ ng_llama.py:157 in forward │\n",
- "│ │\n",
- "│ 154 │ │ self.act_fn = ACT2FN[hidden_act] │\n",
- "│ 155 │ │\n",
- "│ 156 │ def forward(self, x): │\n",
- "│ ❱ 157 │ │ return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x)) │\n",
- "│ 158 │\n",
- "│ 159 │\n",
- "│ 160 class LlamaAttention(nn.Module): │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/torch/nn/modules/module.py:1501 │\n",
- "│ in _call_impl │\n",
- "│ │\n",
- "│ 1498 │ │ if not (self._backward_hooks or self._backward_pre_hooks or self._forward_hooks │\n",
- "│ 1499 │ │ │ │ or _global_backward_pre_hooks or _global_backward_hooks │\n",
- "│ 1500 │ │ │ │ or _global_forward_hooks or _global_forward_pre_hooks): │\n",
- "│ ❱ 1501 │ │ │ return forward_call(*args, **kwargs) │\n",
- "│ 1502 │ │ # Do not call functions when jit is used │\n",
- "│ 1503 │ │ full_backward_hooks, non_full_backward_hooks = [], [] │\n",
- "│ 1504 │ │ backward_pre_hooks = [] │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/accelerate/hooks.py:165 in │\n",
- "│ new_forward │\n",
- "│ │\n",
- "│ 162 │ │ │ with torch.no_grad(): │\n",
- "│ 163 │ │ │ │ output = old_forward(*args, **kwargs) │\n",
- "│ 164 │ │ else: │\n",
- "│ ❱ 165 │ │ │ output = old_forward(*args, **kwargs) │\n",
- "│ 166 │ │ return module._hf_hook.post_forward(module, output) │\n",
- "│ 167 │ │\n",
- "│ 168 │ module.forward = new_forward │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/bitsandbytes/nn/modules.py:320 │\n",
- "│ in forward │\n",
- "│ │\n",
- "│ 317 │ │ if self.bias is not None and self.bias.dtype != x.dtype: │\n",
- "│ 318 │ │ │ self.bias.data = self.bias.data.to(x.dtype) │\n",
- "│ 319 │ │ │\n",
- "│ ❱ 320 │ │ out = bnb.matmul(x, self.weight, bias=self.bias, state=self.state) │\n",
- "│ 321 │ │ │\n",
- "│ 322 │ │ if not self.state.has_fp16_weights: │\n",
- "│ 323 │ │ │ if self.state.CB is not None and self.state.CxB is not None: │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/bitsandbytes/autograd/_functions │\n",
- "│ .py:500 in matmul │\n",
- "│ │\n",
- "│ 497 │ state = state or MatmulLtState() │\n",
- "│ 498 │ if threshold > 0.0: │\n",
- "│ 499 │ │ state.threshold = threshold │\n",
- "│ ❱ 500 │ return MatMul8bitLt.apply(A, B, out, bias, state) │\n",
- "│ 501 │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/torch/autograd/function.py:506 │\n",
- "│ in apply │\n",
- "│ │\n",
- "│ 503 │ │ if not torch._C._are_functorch_transforms_active(): │\n",
- "│ 504 │ │ │ # See NOTE: [functorch vjp and autograd interaction] │\n",
- "│ 505 │ │ │ args = _functorch.utils.unwrap_dead_wrappers(args) │\n",
- "│ ❱ 506 │ │ │ return super().apply(*args, **kwargs) # type: ignore[misc] │\n",
- "│ 507 │ │ │\n",
- "│ 508 │ │ if cls.setup_context == _SingleLevelFunction.setup_context: │\n",
- "│ 509 │ │ │ raise RuntimeError( │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/bitsandbytes/autograd/_functions │\n",
- "│ .py:323 in forward │\n",
- "│ │\n",
- "│ 320 │ │ # 1. Quantize A │\n",
- "│ 321 │ │ if len(A.shape) == 3: │\n",
- "│ 322 │ │ │ A = A.view(-1, A.shape[-1]).contiguous() │\n",
- "│ ❱ 323 │ │ CA, CAt, SCA, SCAt, coo_tensorA = F.double_quant(A.to(torch.float16), threshold= │\n",
- "│ 324 │ │ │\n",
- "│ 325 │ │ if state.threshold > 0.0 and coo_tensorA is not None: │\n",
- "│ 326 │ │ │ if state.has_fp16_weights: │\n",
- "│ │\n",
- "│ /home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/bitsandbytes/functional.py:1660 │\n",
- "│ in double_quant │\n",
- "│ │\n",
- "│ 1657 │ │\n",
- "│ 1658 │ is_on_gpu([A, col_stats, row_stats, out_col, out_row]) │\n",
- "│ 1659 │ if threshold > 0.0: │\n",
- "│ ❱ 1660 │ │ nnz = nnz_row_ptr[-1].item() │\n",
- "│ 1661 │ │ if nnz > 0: │\n",
- "│ 1662 │ │ │ coo_tensor = coo_zeros( │\n",
- "│ 1663 │ │ │ │ A.shape[0], A.shape[1], nnz_row_ptr[-1].item(), device │\n",
- "╰──────────────────────────────────────────────────────────────────────────────────────────────────╯\n",
- "KeyboardInterrupt\n",
- "\n"
- ],
- "text/plain": [
- "\u001b[31m╭─\u001b[0m\u001b[31m──────────────────────────────\u001b[0m\u001b[31m \u001b[0m\u001b[1;31mTraceback \u001b[0m\u001b[1;2;31m(most recent call last)\u001b[0m\u001b[31m \u001b[0m\u001b[31m───────────────────────────────\u001b[0m\u001b[31m─╮\u001b[0m\n",
- "\u001b[31m│\u001b[0m in \u001b[92m\u001b[0m:\u001b[94m94\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m91 \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m92 \u001b[0m\u001b[2m# test\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m93 \u001b[0mdm = IMBDHSDataModule(model, tokenizer) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m94 dm.setup(\u001b[33m'\u001b[0m\u001b[33mtrain\u001b[0m\u001b[33m'\u001b[0m) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m95 \u001b[0mdl = dm.val_dataloader() \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m96 \u001b[0mb = \u001b[96mnext\u001b[0m(\u001b[96miter\u001b[0m(dl)) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m97 \u001b[0mb \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m in \u001b[92msetup\u001b[0m:\u001b[94m39\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m36 \u001b[0m\u001b[2m│ │ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m37 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[96mself\u001b[0m.dataset = load_dataset(\u001b[96mself\u001b[0m.hparams.dataset_name, split=\u001b[33m\"\u001b[0m\u001b[33mtest\u001b[0m\u001b[33m\"\u001b[0m) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m38 \u001b[0m\u001b[2m│ │ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m39 \u001b[2m│ │ \u001b[0mneg_hs, pos_hs, y, all_neg_ans, all_pos_ans = get_hidden_states_many_examples( \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m40 \u001b[0m\u001b[2m│ │ │ \u001b[0m\u001b[96mself\u001b[0m.model, \u001b[96mself\u001b[0m.tokenizer, \u001b[96mself\u001b[0m.dataset, \u001b[96mself\u001b[0m.hparams.model_type, n=\u001b[96mself\u001b[0m.hp \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m41 \u001b[0m\u001b[2m│ │ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m42 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[2m# let's create a simple 50/50 train split (the data is already randomized)\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m in \u001b[92mget_hidden_states_many_examples\u001b[0m:\u001b[94m27\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m24 \u001b[0m\u001b[2m│ │ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m25 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[2m# get hidden states\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m26 \u001b[0m\u001b[2m# print(format_imdb(text, 0))\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m27 \u001b[2m│ │ \u001b[0mneg = get_hidden_states(model, tokenizer, format_imdb(text, \u001b[94m0\u001b[0m), model_type=model \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m28 \u001b[0m\u001b[2m│ │ \u001b[0mpos = get_hidden_states(model, tokenizer, format_imdb(text, \u001b[94m1\u001b[0m), model_type=model \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m29 \u001b[0m\u001b[2m│ │ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m30 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[2m# collect\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m in \u001b[92mget_hidden_states\u001b[0m:\u001b[94m96\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m93 \u001b[0m\u001b[2m# \"encoder\": get_encoder_hidden_states, \"encoder_decoder\": get_encoder_decoder_hidden_st\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m94 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[33m\"\u001b[0m\u001b[33mdecoder\u001b[0m\u001b[33m\"\u001b[0m: get_decoder_hidden_states}[model_type] \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m95 \u001b[0m\u001b[2m│ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m96 \u001b[2m│ \u001b[0m\u001b[94mreturn\u001b[0m fn(model, tokenizer, input_text, layers=layers) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m97 \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m in \u001b[92mget_decoder_hidden_states\u001b[0m:\u001b[94m56\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m53 \u001b[0m\u001b[2m│ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m54 \u001b[0m\u001b[2m│ \u001b[0m\u001b[94mwith\u001b[0m torch.no_grad(): \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m55 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[2m# FIXME: should be a batch, to speed it up\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m56 \u001b[2m│ │ \u001b[0moutput = model(input_ids, \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m57 \u001b[0m\u001b[2m│ │ │ │ │ \u001b[0moutput_hidden_states=\u001b[94mTrue\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m58 \u001b[0m\u001b[2m# , output_attentions=True\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m59 \u001b[0m\u001b[2m│ │ │ │ │ \u001b[0m) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/torch/nn/modules/\u001b[0m\u001b[1;33mmodule.py\u001b[0m:\u001b[94m1501\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m in \u001b[92m_call_impl\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1498 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mif\u001b[0m \u001b[95mnot\u001b[0m (\u001b[96mself\u001b[0m._backward_hooks \u001b[95mor\u001b[0m \u001b[96mself\u001b[0m._backward_pre_hooks \u001b[95mor\u001b[0m \u001b[96mself\u001b[0m._forward_hooks \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1499 \u001b[0m\u001b[2m│ │ │ │ \u001b[0m\u001b[95mor\u001b[0m _global_backward_pre_hooks \u001b[95mor\u001b[0m _global_backward_hooks \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1500 \u001b[0m\u001b[2m│ │ │ │ \u001b[0m\u001b[95mor\u001b[0m _global_forward_hooks \u001b[95mor\u001b[0m _global_forward_pre_hooks): \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m1501 \u001b[2m│ │ │ \u001b[0m\u001b[94mreturn\u001b[0m forward_call(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1502 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[2m# Do not call functions when jit is used\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1503 \u001b[0m\u001b[2m│ │ \u001b[0mfull_backward_hooks, non_full_backward_hooks = [], [] \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1504 \u001b[0m\u001b[2m│ │ \u001b[0mbackward_pre_hooks = [] \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/accelerate/\u001b[0m\u001b[1;33mhooks.py\u001b[0m:\u001b[94m165\u001b[0m in \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[92mnew_forward\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m162 \u001b[0m\u001b[2m│ │ │ \u001b[0m\u001b[94mwith\u001b[0m torch.no_grad(): \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m163 \u001b[0m\u001b[2m│ │ │ │ \u001b[0moutput = old_forward(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m164 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94melse\u001b[0m: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m165 \u001b[2m│ │ │ \u001b[0moutput = old_forward(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m166 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mreturn\u001b[0m module._hf_hook.post_forward(module, output) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m167 \u001b[0m\u001b[2m│ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m168 \u001b[0m\u001b[2m│ \u001b[0mmodule.forward = new_forward \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/transformers/models/llama/\u001b[0m\u001b[1;33mmodeli\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[1;33mng_llama.py\u001b[0m:\u001b[94m687\u001b[0m in \u001b[92mforward\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m684 \u001b[0m\u001b[2m│ │ \u001b[0mreturn_dict = return_dict \u001b[94mif\u001b[0m return_dict \u001b[95mis\u001b[0m \u001b[95mnot\u001b[0m \u001b[94mNone\u001b[0m \u001b[94melse\u001b[0m \u001b[96mself\u001b[0m.config.use_return \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m685 \u001b[0m\u001b[2m│ │ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m686 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[2m# decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m687 \u001b[2m│ │ \u001b[0moutputs = \u001b[96mself\u001b[0m.model( \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m688 \u001b[0m\u001b[2m│ │ │ \u001b[0minput_ids=input_ids, \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m689 \u001b[0m\u001b[2m│ │ │ \u001b[0mattention_mask=attention_mask, \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m690 \u001b[0m\u001b[2m│ │ │ \u001b[0mposition_ids=position_ids, \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/torch/nn/modules/\u001b[0m\u001b[1;33mmodule.py\u001b[0m:\u001b[94m1501\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m in \u001b[92m_call_impl\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1498 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mif\u001b[0m \u001b[95mnot\u001b[0m (\u001b[96mself\u001b[0m._backward_hooks \u001b[95mor\u001b[0m \u001b[96mself\u001b[0m._backward_pre_hooks \u001b[95mor\u001b[0m \u001b[96mself\u001b[0m._forward_hooks \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1499 \u001b[0m\u001b[2m│ │ │ │ \u001b[0m\u001b[95mor\u001b[0m _global_backward_pre_hooks \u001b[95mor\u001b[0m _global_backward_hooks \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1500 \u001b[0m\u001b[2m│ │ │ │ \u001b[0m\u001b[95mor\u001b[0m _global_forward_hooks \u001b[95mor\u001b[0m _global_forward_pre_hooks): \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m1501 \u001b[2m│ │ │ \u001b[0m\u001b[94mreturn\u001b[0m forward_call(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1502 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[2m# Do not call functions when jit is used\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1503 \u001b[0m\u001b[2m│ │ \u001b[0mfull_backward_hooks, non_full_backward_hooks = [], [] \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1504 \u001b[0m\u001b[2m│ │ \u001b[0mbackward_pre_hooks = [] \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/accelerate/\u001b[0m\u001b[1;33mhooks.py\u001b[0m:\u001b[94m165\u001b[0m in \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[92mnew_forward\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m162 \u001b[0m\u001b[2m│ │ │ \u001b[0m\u001b[94mwith\u001b[0m torch.no_grad(): \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m163 \u001b[0m\u001b[2m│ │ │ │ \u001b[0moutput = old_forward(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m164 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94melse\u001b[0m: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m165 \u001b[2m│ │ │ \u001b[0moutput = old_forward(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m166 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mreturn\u001b[0m module._hf_hook.post_forward(module, output) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m167 \u001b[0m\u001b[2m│ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m168 \u001b[0m\u001b[2m│ \u001b[0mmodule.forward = new_forward \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/transformers/models/llama/\u001b[0m\u001b[1;33mmodeli\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[1;33mng_llama.py\u001b[0m:\u001b[94m577\u001b[0m in \u001b[92mforward\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m574 \u001b[0m\u001b[2m│ │ │ │ │ \u001b[0m\u001b[94mNone\u001b[0m, \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m575 \u001b[0m\u001b[2m│ │ │ │ \u001b[0m) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m576 \u001b[0m\u001b[2m│ │ │ \u001b[0m\u001b[94melse\u001b[0m: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m577 \u001b[2m│ │ │ │ \u001b[0mlayer_outputs = decoder_layer( \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m578 \u001b[0m\u001b[2m│ │ │ │ │ \u001b[0mhidden_states, \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m579 \u001b[0m\u001b[2m│ │ │ │ │ \u001b[0mattention_mask=attention_mask, \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m580 \u001b[0m\u001b[2m│ │ │ │ │ \u001b[0mposition_ids=position_ids, \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/torch/nn/modules/\u001b[0m\u001b[1;33mmodule.py\u001b[0m:\u001b[94m1501\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m in \u001b[92m_call_impl\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1498 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mif\u001b[0m \u001b[95mnot\u001b[0m (\u001b[96mself\u001b[0m._backward_hooks \u001b[95mor\u001b[0m \u001b[96mself\u001b[0m._backward_pre_hooks \u001b[95mor\u001b[0m \u001b[96mself\u001b[0m._forward_hooks \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1499 \u001b[0m\u001b[2m│ │ │ │ \u001b[0m\u001b[95mor\u001b[0m _global_backward_pre_hooks \u001b[95mor\u001b[0m _global_backward_hooks \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1500 \u001b[0m\u001b[2m│ │ │ │ \u001b[0m\u001b[95mor\u001b[0m _global_forward_hooks \u001b[95mor\u001b[0m _global_forward_pre_hooks): \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m1501 \u001b[2m│ │ │ \u001b[0m\u001b[94mreturn\u001b[0m forward_call(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1502 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[2m# Do not call functions when jit is used\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1503 \u001b[0m\u001b[2m│ │ \u001b[0mfull_backward_hooks, non_full_backward_hooks = [], [] \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1504 \u001b[0m\u001b[2m│ │ \u001b[0mbackward_pre_hooks = [] \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/accelerate/\u001b[0m\u001b[1;33mhooks.py\u001b[0m:\u001b[94m165\u001b[0m in \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[92mnew_forward\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m162 \u001b[0m\u001b[2m│ │ │ \u001b[0m\u001b[94mwith\u001b[0m torch.no_grad(): \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m163 \u001b[0m\u001b[2m│ │ │ │ \u001b[0moutput = old_forward(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m164 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94melse\u001b[0m: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m165 \u001b[2m│ │ │ \u001b[0moutput = old_forward(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m166 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mreturn\u001b[0m module._hf_hook.post_forward(module, output) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m167 \u001b[0m\u001b[2m│ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m168 \u001b[0m\u001b[2m│ \u001b[0mmodule.forward = new_forward \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/transformers/models/llama/\u001b[0m\u001b[1;33mmodeli\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[1;33mng_llama.py\u001b[0m:\u001b[94m305\u001b[0m in \u001b[92mforward\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m302 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[2m# Fully Connected\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m303 \u001b[0m\u001b[2m│ │ \u001b[0mresidual = hidden_states \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m304 \u001b[0m\u001b[2m│ │ \u001b[0mhidden_states = \u001b[96mself\u001b[0m.post_attention_layernorm(hidden_states) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m305 \u001b[2m│ │ \u001b[0mhidden_states = \u001b[96mself\u001b[0m.mlp(hidden_states) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m306 \u001b[0m\u001b[2m│ │ \u001b[0mhidden_states = residual + hidden_states \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m307 \u001b[0m\u001b[2m│ │ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m308 \u001b[0m\u001b[2m│ │ \u001b[0moutputs = (hidden_states,) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/torch/nn/modules/\u001b[0m\u001b[1;33mmodule.py\u001b[0m:\u001b[94m1501\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m in \u001b[92m_call_impl\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1498 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mif\u001b[0m \u001b[95mnot\u001b[0m (\u001b[96mself\u001b[0m._backward_hooks \u001b[95mor\u001b[0m \u001b[96mself\u001b[0m._backward_pre_hooks \u001b[95mor\u001b[0m \u001b[96mself\u001b[0m._forward_hooks \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1499 \u001b[0m\u001b[2m│ │ │ │ \u001b[0m\u001b[95mor\u001b[0m _global_backward_pre_hooks \u001b[95mor\u001b[0m _global_backward_hooks \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1500 \u001b[0m\u001b[2m│ │ │ │ \u001b[0m\u001b[95mor\u001b[0m _global_forward_hooks \u001b[95mor\u001b[0m _global_forward_pre_hooks): \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m1501 \u001b[2m│ │ │ \u001b[0m\u001b[94mreturn\u001b[0m forward_call(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1502 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[2m# Do not call functions when jit is used\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1503 \u001b[0m\u001b[2m│ │ \u001b[0mfull_backward_hooks, non_full_backward_hooks = [], [] \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1504 \u001b[0m\u001b[2m│ │ \u001b[0mbackward_pre_hooks = [] \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/accelerate/\u001b[0m\u001b[1;33mhooks.py\u001b[0m:\u001b[94m165\u001b[0m in \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[92mnew_forward\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m162 \u001b[0m\u001b[2m│ │ │ \u001b[0m\u001b[94mwith\u001b[0m torch.no_grad(): \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m163 \u001b[0m\u001b[2m│ │ │ │ \u001b[0moutput = old_forward(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m164 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94melse\u001b[0m: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m165 \u001b[2m│ │ │ \u001b[0moutput = old_forward(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m166 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mreturn\u001b[0m module._hf_hook.post_forward(module, output) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m167 \u001b[0m\u001b[2m│ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m168 \u001b[0m\u001b[2m│ \u001b[0mmodule.forward = new_forward \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/transformers/models/llama/\u001b[0m\u001b[1;33mmodeli\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[1;33mng_llama.py\u001b[0m:\u001b[94m157\u001b[0m in \u001b[92mforward\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m154 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[96mself\u001b[0m.act_fn = ACT2FN[hidden_act] \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m155 \u001b[0m\u001b[2m│ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m156 \u001b[0m\u001b[2m│ \u001b[0m\u001b[94mdef\u001b[0m \u001b[92mforward\u001b[0m(\u001b[96mself\u001b[0m, x): \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m157 \u001b[2m│ │ \u001b[0m\u001b[94mreturn\u001b[0m \u001b[96mself\u001b[0m.down_proj(\u001b[96mself\u001b[0m.act_fn(\u001b[96mself\u001b[0m.gate_proj(x)) * \u001b[96mself\u001b[0m.up_proj(x)) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m158 \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m159 \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m160 \u001b[0m\u001b[94mclass\u001b[0m \u001b[4;92mLlamaAttention\u001b[0m(nn.Module): \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/torch/nn/modules/\u001b[0m\u001b[1;33mmodule.py\u001b[0m:\u001b[94m1501\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m in \u001b[92m_call_impl\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1498 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mif\u001b[0m \u001b[95mnot\u001b[0m (\u001b[96mself\u001b[0m._backward_hooks \u001b[95mor\u001b[0m \u001b[96mself\u001b[0m._backward_pre_hooks \u001b[95mor\u001b[0m \u001b[96mself\u001b[0m._forward_hooks \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1499 \u001b[0m\u001b[2m│ │ │ │ \u001b[0m\u001b[95mor\u001b[0m _global_backward_pre_hooks \u001b[95mor\u001b[0m _global_backward_hooks \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1500 \u001b[0m\u001b[2m│ │ │ │ \u001b[0m\u001b[95mor\u001b[0m _global_forward_hooks \u001b[95mor\u001b[0m _global_forward_pre_hooks): \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m1501 \u001b[2m│ │ │ \u001b[0m\u001b[94mreturn\u001b[0m forward_call(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1502 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[2m# Do not call functions when jit is used\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1503 \u001b[0m\u001b[2m│ │ \u001b[0mfull_backward_hooks, non_full_backward_hooks = [], [] \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1504 \u001b[0m\u001b[2m│ │ \u001b[0mbackward_pre_hooks = [] \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/accelerate/\u001b[0m\u001b[1;33mhooks.py\u001b[0m:\u001b[94m165\u001b[0m in \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[92mnew_forward\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m162 \u001b[0m\u001b[2m│ │ │ \u001b[0m\u001b[94mwith\u001b[0m torch.no_grad(): \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m163 \u001b[0m\u001b[2m│ │ │ │ \u001b[0moutput = old_forward(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m164 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94melse\u001b[0m: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m165 \u001b[2m│ │ │ \u001b[0moutput = old_forward(*args, **kwargs) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m166 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mreturn\u001b[0m module._hf_hook.post_forward(module, output) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m167 \u001b[0m\u001b[2m│ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m168 \u001b[0m\u001b[2m│ \u001b[0mmodule.forward = new_forward \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/bitsandbytes/nn/\u001b[0m\u001b[1;33mmodules.py\u001b[0m:\u001b[94m320\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m in \u001b[92mforward\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m317 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mif\u001b[0m \u001b[96mself\u001b[0m.bias \u001b[95mis\u001b[0m \u001b[95mnot\u001b[0m \u001b[94mNone\u001b[0m \u001b[95mand\u001b[0m \u001b[96mself\u001b[0m.bias.dtype != x.dtype: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m318 \u001b[0m\u001b[2m│ │ │ \u001b[0m\u001b[96mself\u001b[0m.bias.data = \u001b[96mself\u001b[0m.bias.data.to(x.dtype) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m319 \u001b[0m\u001b[2m│ │ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m320 \u001b[2m│ │ \u001b[0mout = bnb.matmul(x, \u001b[96mself\u001b[0m.weight, bias=\u001b[96mself\u001b[0m.bias, state=\u001b[96mself\u001b[0m.state) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m321 \u001b[0m\u001b[2m│ │ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m322 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mif\u001b[0m \u001b[95mnot\u001b[0m \u001b[96mself\u001b[0m.state.has_fp16_weights: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m323 \u001b[0m\u001b[2m│ │ │ \u001b[0m\u001b[94mif\u001b[0m \u001b[96mself\u001b[0m.state.CB \u001b[95mis\u001b[0m \u001b[95mnot\u001b[0m \u001b[94mNone\u001b[0m \u001b[95mand\u001b[0m \u001b[96mself\u001b[0m.state.CxB \u001b[95mis\u001b[0m \u001b[95mnot\u001b[0m \u001b[94mNone\u001b[0m: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/bitsandbytes/autograd/\u001b[0m\u001b[1;33m_functions\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[1;33m.py\u001b[0m:\u001b[94m500\u001b[0m in \u001b[92mmatmul\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m497 \u001b[0m\u001b[2m│ \u001b[0mstate = state \u001b[95mor\u001b[0m MatmulLtState() \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m498 \u001b[0m\u001b[2m│ \u001b[0m\u001b[94mif\u001b[0m threshold > \u001b[94m0.0\u001b[0m: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m499 \u001b[0m\u001b[2m│ │ \u001b[0mstate.threshold = threshold \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m500 \u001b[2m│ \u001b[0m\u001b[94mreturn\u001b[0m MatMul8bitLt.apply(A, B, out, bias, state) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m501 \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/torch/autograd/\u001b[0m\u001b[1;33mfunction.py\u001b[0m:\u001b[94m506\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m in \u001b[92mapply\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m503 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mif\u001b[0m \u001b[95mnot\u001b[0m torch._C._are_functorch_transforms_active(): \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m504 \u001b[0m\u001b[2m│ │ │ \u001b[0m\u001b[2m# See NOTE: [functorch vjp and autograd interaction]\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m505 \u001b[0m\u001b[2m│ │ │ \u001b[0margs = _functorch.utils.unwrap_dead_wrappers(args) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m506 \u001b[2m│ │ │ \u001b[0m\u001b[94mreturn\u001b[0m \u001b[96msuper\u001b[0m().apply(*args, **kwargs) \u001b[2m# type: ignore[misc]\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m507 \u001b[0m\u001b[2m│ │ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m508 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mif\u001b[0m \u001b[96mcls\u001b[0m.setup_context == _SingleLevelFunction.setup_context: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m509 \u001b[0m\u001b[2m│ │ │ \u001b[0m\u001b[94mraise\u001b[0m \u001b[96mRuntimeError\u001b[0m( \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/bitsandbytes/autograd/\u001b[0m\u001b[1;33m_functions\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[1;33m.py\u001b[0m:\u001b[94m323\u001b[0m in \u001b[92mforward\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m320 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[2m# 1. Quantize A\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m321 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mif\u001b[0m \u001b[96mlen\u001b[0m(A.shape) == \u001b[94m3\u001b[0m: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m322 \u001b[0m\u001b[2m│ │ │ \u001b[0mA = A.view(-\u001b[94m1\u001b[0m, A.shape[-\u001b[94m1\u001b[0m]).contiguous() \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m323 \u001b[2m│ │ \u001b[0mCA, CAt, SCA, SCAt, coo_tensorA = F.double_quant(A.to(torch.float16), threshold= \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m324 \u001b[0m\u001b[2m│ │ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m325 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mif\u001b[0m state.threshold > \u001b[94m0.0\u001b[0m \u001b[95mand\u001b[0m coo_tensorA \u001b[95mis\u001b[0m \u001b[95mnot\u001b[0m \u001b[94mNone\u001b[0m: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m326 \u001b[0m\u001b[2m│ │ │ \u001b[0m\u001b[94mif\u001b[0m state.has_fp16_weights: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2;33m/home/wassname/miniforge3/envs/dlk2/lib/python3.9/site-packages/bitsandbytes/\u001b[0m\u001b[1;33mfunctional.py\u001b[0m:\u001b[94m1660\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m in \u001b[92mdouble_quant\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1657 \u001b[0m\u001b[2m│ \u001b[0m \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1658 \u001b[0m\u001b[2m│ \u001b[0mis_on_gpu([A, col_stats, row_stats, out_col, out_row]) \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1659 \u001b[0m\u001b[2m│ \u001b[0m\u001b[94mif\u001b[0m threshold > \u001b[94m0.0\u001b[0m: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[31m❱ \u001b[0m1660 \u001b[2m│ │ \u001b[0mnnz = nnz_row_ptr[-\u001b[94m1\u001b[0m].item() \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1661 \u001b[0m\u001b[2m│ │ \u001b[0m\u001b[94mif\u001b[0m nnz > \u001b[94m0\u001b[0m: \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1662 \u001b[0m\u001b[2m│ │ │ \u001b[0mcoo_tensor = coo_zeros( \u001b[31m│\u001b[0m\n",
- "\u001b[31m│\u001b[0m \u001b[2m1663 \u001b[0m\u001b[2m│ │ │ │ \u001b[0mA.shape[\u001b[94m0\u001b[0m], A.shape[\u001b[94m1\u001b[0m], nnz_row_ptr[-\u001b[94m1\u001b[0m].item(), device \u001b[31m│\u001b[0m\n",
- "\u001b[31m╰──────────────────────────────────────────────────────────────────────────────────────────────────╯\u001b[0m\n",
- "\u001b[1;91mKeyboardInterrupt\u001b[0m\n"
- ]
- },
- "metadata": {},
- "output_type": "display_data"
- }
- ],
+ "outputs": [],
"source": [
- "from dataclasses import dataclass\n",
- "from torch.utils.data import random_split, DataLoader, TensorDataset\n",
- "from transformers.models.auto.modeling_auto import AutoModel\n",
- "# from scipy.stats import zscore\n",
- "\n",
- "from sklearn.preprocessing import RobustScaler\n",
"\n",
"# def normalize(x):\n",
"# \"\"\"\n",
@@ -1688,8 +1486,8 @@
"execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T01:18:07.508827Z",
- "start_time": "2023-05-19T01:18:07.508819Z"
+ "end_time": "2023-05-20T01:57:03.647107Z",
+ "start_time": "2023-05-20T01:57:03.647101Z"
}
},
"outputs": [],
@@ -1709,8 +1507,8 @@
"execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T01:18:07.509553Z",
- "start_time": "2023-05-19T01:18:07.509545Z"
+ "end_time": "2023-05-20T01:57:03.647739Z",
+ "start_time": "2023-05-20T01:57:03.647733Z"
}
},
"outputs": [],
@@ -1723,8 +1521,8 @@
"execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T01:18:07.510256Z",
- "start_time": "2023-05-19T01:18:07.510249Z"
+ "end_time": "2023-05-20T01:57:03.648618Z",
+ "start_time": "2023-05-20T01:57:03.648611Z"
}
},
"outputs": [],
@@ -1835,8 +1633,8 @@
"execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T01:18:07.510853Z",
- "start_time": "2023-05-19T01:18:07.510846Z"
+ "end_time": "2023-05-20T01:57:03.649120Z",
+ "start_time": "2023-05-20T01:57:03.649114Z"
}
},
"outputs": [],
@@ -1852,8 +1650,8 @@
"execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T01:18:07.511523Z",
- "start_time": "2023-05-19T01:18:07.511516Z"
+ "end_time": "2023-05-20T01:57:03.649727Z",
+ "start_time": "2023-05-20T01:57:03.649721Z"
}
},
"outputs": [],
@@ -1866,8 +1664,8 @@
"execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T01:18:07.512087Z",
- "start_time": "2023-05-19T01:18:07.512080Z"
+ "end_time": "2023-05-20T01:57:03.650409Z",
+ "start_time": "2023-05-20T01:57:03.650402Z"
},
"scrolled": true
},
@@ -1883,8 +1681,8 @@
"execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T01:18:07.512749Z",
- "start_time": "2023-05-19T01:18:07.512742Z"
+ "end_time": "2023-05-20T01:57:03.650895Z",
+ "start_time": "2023-05-20T01:57:03.650889Z"
}
},
"outputs": [],
@@ -1909,8 +1707,8 @@
"execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T01:18:07.513235Z",
- "start_time": "2023-05-19T01:18:07.513229Z"
+ "end_time": "2023-05-20T01:57:03.651740Z",
+ "start_time": "2023-05-20T01:57:03.651734Z"
}
},
"outputs": [],
@@ -1946,8 +1744,8 @@
"execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T01:18:07.513826Z",
- "start_time": "2023-05-19T01:18:07.513819Z"
+ "end_time": "2023-05-20T01:57:03.652269Z",
+ "start_time": "2023-05-20T01:57:03.652263Z"
}
},
"outputs": [],
@@ -1961,8 +1759,8 @@
"execution_count": null,
"metadata": {
"ExecuteTime": {
- "end_time": "2023-05-19T01:18:07.514335Z",
- "start_time": "2023-05-19T01:18:07.514329Z"
+ "end_time": "2023-05-20T01:57:03.653040Z",
+ "start_time": "2023-05-20T01:57:03.653033Z"
}
},
"outputs": [],
diff --git a/mjc_notes.md b/mjc_notes.md
index d719443..7cacf9d 100644
--- a/mjc_notes.md
+++ b/mjc_notes.md
@@ -11,4 +11,5 @@ pip install -r requirements.txt
- [x] Convert it to lightning
- [ ] batch for get hidden states
- - [ ] and cache
+ - [x] and cache
+ - [ ] 9s vs 60. so 10x faster
|