{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"def read_data(data):\n    return tuple(d for d in data[:-1]), data[-1]","metadata":{"_uuid":"c9429fef-ac71-440d-be54-44fbad509564","_cell_guid":"55095601-efb8-4ac0-b943-3ec3fda07160","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:46:41.883566Z","iopub.execute_input":"2022-07-21T13:46:41.883910Z","iopub.status.idle":"2022-07-21T13:46:41.910402Z","shell.execute_reply.started":"2022-07-21T13:46:41.883836Z","shell.execute_reply":"2022-07-21T13:46:41.909596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#import wandb\n#wandb.login()","metadata":{"_uuid":"1003a634-97ef-4819-af68-6f68788294b4","_cell_guid":"b0d3e4d8-a3ce-4cf3-b9b9-5abf999efa5e","collapsed":false,"_kg_hide-input":true,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T08:43:12.593861Z","iopub.execute_input":"2022-07-21T08:43:12.594447Z","iopub.status.idle":"2022-07-21T08:43:12.598557Z","shell.execute_reply.started":"2022-07-21T08:43:12.594411Z","shell.execute_reply":"2022-07-21T08:43:12.597683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#wandb.init(project=\"ai4codeT5\")","metadata":{"_uuid":"8186e686-2d58-4946-9e64-7e0547314e04","_cell_guid":"c211cfd1-34c9-4f7d-bdd0-9d67c322d115","collapsed":false,"_kg_hide-input":true,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T08:43:13.188759Z","iopub.execute_input":"2022-07-21T08:43:13.189485Z","iopub.status.idle":"2022-07-21T08:43:13.193751Z","shell.execute_reply.started":"2022-07-21T08:43:13.189449Z","shell.execute_reply":"2022-07-21T08:43:13.192625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''wandb.config = {\n  \"learning_rate\": 0.001,\n  \"epochs\": 10,\n  \"batch_size\": 4\n}'''","metadata":{"_uuid":"e7ebebe3-1358-4133-ad71-7944e2f3a2b7","_cell_guid":"b7de60fd-0138-4543-86c1-0e87bd0f87d5","collapsed":false,"_kg_hide-input":true,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T08:43:13.401274Z","iopub.execute_input":"2022-07-21T08:43:13.401621Z","iopub.status.idle":"2022-07-21T08:43:13.410220Z","shell.execute_reply.started":"2022-07-21T08:43:13.401592Z","shell.execute_reply":"2022-07-21T08:43:13.409468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Models.py\nimport torch\nimport torch.nn as nn\nimport numpy as np\nfrom transformers import (RobertaConfig, RobertaModel, RobertaTokenizer,\n                          BartConfig, BartForConditionalGeneration, BartTokenizer,\n                          T5Config, T5ForConditionalGeneration, T5Tokenizer)\nimport logging\n\nlogger = logging.getLogger(__name__)\n\nMODEL_CLASSES = {\n                 \n                 'codet5': (T5Config, T5ForConditionalGeneration, RobertaTokenizer)}","metadata":{"_uuid":"608bae6b-b8e0-4284-b98f-2c05831373f3","_cell_guid":"4ca9162e-3da8-49bd-8e43-f03a29ffb035","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:46:41.913033Z","iopub.execute_input":"2022-07-21T13:46:41.914852Z","iopub.status.idle":"2022-07-21T13:46:49.340743Z","shell.execute_reply.started":"2022-07-21T13:46:41.914821Z","shell.execute_reply":"2022-07-21T13:46:49.339930Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODEL_CLASSES['codet5']","metadata":{"_uuid":"5610cbee-e47d-4b3a-9a84-4b190569b387","_cell_guid":"20e71ab0-251c-46ab-895f-d0fca8b173f1","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:46:49.342301Z","iopub.execute_input":"2022-07-21T13:46:49.342907Z","iopub.status.idle":"2022-07-21T13:46:49.349720Z","shell.execute_reply.started":"2022-07-21T13:46:49.342870Z","shell.execute_reply":"2022-07-21T13:46:49.349027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model_size(model):\n    model_parameters = filter(lambda p: p.requires_grad, model.parameters())\n    model_size = sum([np.prod(p.size()) for p in model_parameters])\n    return \"{}M\".format(round(model_size / 1e+6))","metadata":{"_uuid":"8dd5c778-8a5b-4108-95aa-a32c110e7ddd","_cell_guid":"17af2d19-da65-4896-8716-83efb3c72f20","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:46:49.350939Z","iopub.execute_input":"2022-07-21T13:46:49.351508Z","iopub.status.idle":"2022-07-21T13:46:49.361127Z","shell.execute_reply.started":"2022-07-21T13:46:49.351471Z","shell.execute_reply":"2022-07-21T13:46:49.360342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\nimport torch\nimport logging\nimport multiprocessing\nimport numpy as np\nfrom torch.utils.data import DataLoader\nlogger = logging.getLogger(__name__)","metadata":{"_uuid":"84c395b5-ae46-46ab-93a6-1b0f31f03d9b","_cell_guid":"78fe78c4-bba8-437f-be60-4c717d2c2ce3","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:46:49.363176Z","iopub.execute_input":"2022-07-21T13:46:49.363584Z","iopub.status.idle":"2022-07-21T13:46:49.371970Z","shell.execute_reply.started":"2022-07-21T13:46:49.363548Z","shell.execute_reply":"2022-07-21T13:46:49.371124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_dist():\n    # Setup CUDA, GPU & distributed training\n   \n        # Setup for distributed data parallel\n    torch.cuda.set_device(0)\n    device = torch.device(\"cuda\", 0)\n    torch.distributed.init_process_group(backend='nccl')\n    #args.n_gpu = 1\n    cpu_cont = multiprocessing.cpu_count()\n    logger.warning(\"Process rank: %s, device: %s, n_gpu: %s, distributed training: %s, cpu count: %d\",\n                   0, device, 1, bool(0!= -1), cpu_cont)\n    #args.device = device\n    #args.cpu_cont = cpu_cont\n\n\ndef set_seed():\n    \"\"\"set random seed.\"\"\"\n    random.seed(1234)\n    np.random.seed(1234)\n    torch.manual_seed(1234)\n    if 1 > 0:\n        torch.cuda.manual_seed_all(1234)","metadata":{"_uuid":"865f7ad9-2466-4aba-a7e1-1bbfddb7b77d","_cell_guid":"faf21c3b-8389-4bd9-9d20-90eab7f6669e","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:46:49.373074Z","iopub.execute_input":"2022-07-21T13:46:49.373466Z","iopub.status.idle":"2022-07-21T13:46:49.383720Z","shell.execute_reply.started":"2022-07-21T13:46:49.373431Z","shell.execute_reply":"2022-07-21T13:46:49.382973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def build_or_load_gen_model():\n    config_class, model_class, tokenizer_class = MODEL_CLASSES['codet5']\n    config = config_class.from_pretrained(\"Salesforce/codet5-small\")\n    tokenizer = tokenizer_class.from_pretrained(\"Salesforce/codet5-small\")\n    \n    model = model_class.from_pretrained(\"Salesforce/codet5-small\")\n\n    logger.info(\"Finish loading model [%s] from %s\", get_model_size(model), \"Salesforce/codet5-small\")\n\n    #if load_model_path is not None:\n        #logger.info(\"Reload model from {}\".format(load_model_path))\n        #model.load_state_dict(torch.load(load_model_path))\n\n    return config, model, tokenizer","metadata":{"_uuid":"d37f4405-90e2-467a-a126-a61cf55d1372","_cell_guid":"db3d2bd7-8463-42f9-90fa-1de9c53cd63e","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:46:49.385380Z","iopub.execute_input":"2022-07-21T13:46:49.385694Z","iopub.status.idle":"2022-07-21T13:46:49.395254Z","shell.execute_reply.started":"2022-07-21T13:46:49.385670Z","shell.execute_reply":"2022-07-21T13:46:49.394539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import nn\nclass RobertaClassificationHead(nn.Module):\n    \"\"\"Head for sentence-level classification tasks.\"\"\"\n\n    def __init__(self, config):\n        super().__init__()\n        self.dense = nn.Linear(config.hidden_size * 2, config.hidden_size)\n        self.out_proj = nn.Linear(config.hidden_size, 2)\n\n    def forward(self, x, **kwargs):\n        x = x.reshape(-1, x.size(-1) * 2)\n        x = self.dense(x)\n        x = torch.tanh(x)\n        x = self.out_proj(x)\n        return x","metadata":{"_uuid":"99a87ee1-0f31-43af-9ae4-e5b82ef58149","_cell_guid":"165d4c35-e9c3-4a3f-b5da-999e9e80cafd","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:47:11.753998Z","iopub.execute_input":"2022-07-21T13:47:11.754721Z","iopub.status.idle":"2022-07-21T13:47:11.761238Z","shell.execute_reply.started":"2022-07-21T13:47:11.754689Z","shell.execute_reply":"2022-07-21T13:47:11.760015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def softXEnt (input, target):\n        logprobs = torch.nn.functional.log_softmax (input, dim = 1)\n        return  -(target * logprobs).sum() / input.shape[0]","metadata":{"_uuid":"ec4c3117-4794-414e-8d75-e86cb73b2f6a","_cell_guid":"b2a6efff-3ee0-4144-8075-a08404fef0d6","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:47:12.264600Z","iopub.execute_input":"2022-07-21T13:47:12.264937Z","iopub.status.idle":"2022-07-21T13:47:12.270003Z","shell.execute_reply.started":"2022-07-21T13:47:12.264908Z","shell.execute_reply":"2022-07-21T13:47:12.269136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SortModel(nn.Module):\n    def __init__(self, encoder, config,tokenizer):\n        super(SortModel, self).__init__()\n        self.encoder = encoder\n        self.config = config\n        self.tokenizer = tokenizer\n        self.classifier = nn.Linear(config.hidden_size, 1)\n\n    def get_t5_vec(self, source_ids):\n        attention_mask = source_ids.ne(self.tokenizer.pad_token_id)\n        outputs = self.encoder(input_ids=source_ids, attention_mask=attention_mask,\n                               labels=source_ids, decoder_attention_mask=attention_mask, output_hidden_states=True)\n        hidden_states = outputs['decoder_hidden_states'][-1]\n        eos_mask = source_ids.eq(self.config.eos_token_id)\n\n        if len(torch.unique(eos_mask.sum(1))) > 1:\n            raise ValueError(\"All examples must have the same number of <eos> tokens.\")\n        vec = hidden_states[eos_mask, :].view(hidden_states.size(0), -1,\n                                              hidden_states.size(-1))[:, -1, :]\n        return vec\n    \n    def softXEnt (input, target):\n        logprobs = torch.nn.functional.log_softmax (input, dim = 1)\n        return  -(target * logprobs).sum() / input.shape[0]\n    \n\n    def forward(self,input_ids=None, labels=None,fts = None):\n        max_src_length=512\n        input_ids = input_ids.view(-1,max_src_length )\n\n        vec = self.get_t5_vec(input_ids)\n        #logits = torch.cat((vec,fts),1)\n        logits = self.classifier(vec) \n        prob = nn.functional.softmax(logits,dim=0)\n        if labels is not None:\n            #loss = nn.CrossEntropyLoss()\n            loss = torch.nn.L1Loss()\n            output = loss(logits,labels)\n            return output,logits\n        else:\n            return prob","metadata":{"_uuid":"45cc3911-a95c-4da5-82fe-8ce13295b25c","_cell_guid":"912f0b8a-4bcd-48d8-8fd1-232e2af3fd7f","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:47:12.749763Z","iopub.execute_input":"2022-07-21T13:47:12.750430Z","iopub.status.idle":"2022-07-21T13:47:12.763120Z","shell.execute_reply.started":"2022-07-21T13:47:12.750394Z","shell.execute_reply":"2022-07-21T13:47:12.761590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\n\nclass Example(object):\n    \"\"\"A single training/test example.\"\"\"\n\n    def __init__(self,\n                 source,\n                 idx=None,\n                 target=None,\n                 ):\n        self.id = idx\n        self.source = source\n        self.target = target\n        \n    \n    #def __getitem__(self,index) : \n    \n    def __len__(self):\n        return self.source.shape[0]\nclass InputFeatures(object):\n    \"\"\"A single training/test features for a example.\"\"\"\n\n    def __init__(self,\n                 example_id,\n                 source_ids,\n                 target_ids,\n                 fts\n                 ):\n        self.example_id = example_id\n        self.source_ids = source_ids\n        self.target_ids = target_ids\n        self.fts=fts","metadata":{"_uuid":"2d0900da-7ee4-48f2-94ed-c8bfed4e4ff8","_cell_guid":"6ee91ed2-1ed7-4b0a-86b5-207c1cdea28a","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:47:14.538758Z","iopub.execute_input":"2022-07-21T13:47:14.539151Z","iopub.status.idle":"2022-07-21T13:47:14.545916Z","shell.execute_reply.started":"2022-07-21T13:47:14.539118Z","shell.execute_reply":"2022-07-21T13:47:14.545023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import TensorDataset\nimport numpy as np\nimport logging\nimport os\nimport random\nimport torch\nimport time\nfrom tqdm import tqdm\n\nlogger = logging.getLogger(__name__)\n\n\ndef load_and_cache_gen_data(data_num,filename, tokenizer, split_tag, only_src=False, is_sample=False):\n    # cache the data into args.cache_path except it is sampled\n    # only_src: control whether to return only source ids for bleu evaluating (dev/test)\n    # return: examples (Example object), data (TensorDataset)\n    data_tag = '_all' if data_num == -1 else '_%d' % data_num\n    cache_fn = '{}/{}.pt'.format('./outputs', split_tag + ('_src' if only_src else '') + data_tag)\n\n    examples,df = read_sort_examples(filename, data_num)\n\n    if is_sample:\n        examples = random.sample(examples, min(5000, len(examples)))\n    calc_stats(examples, tokenizer, is_tokenize=True)\n    if os.path.exists(cache_fn) and not is_sample:\n        logger.info(\"Load cache data from %s\", cache_fn)\n        data = torch.load(cache_fn)\n    else:\n        if is_sample:\n            logger.info(\"Sample 5k data for computing bleu from %s\", filename)\n        else:\n            logger.info(\"Create cache data into %s\", cache_fn)\n        tuple_examples = [(example, idx, tokenizer,fts,512) for idx, example in enumerate(examples)]\n        features = map(convert_examples_to_features, tqdm(tuple_examples, total=len(tuple_examples)))\n        featuresl = list(features)\n        all_source_ids = torch.tensor([f.source_ids for f in featuresl], dtype=torch.long)\n        all_target_ids = torch.tensor([f.target_ids for f in featuresl], dtype=torch.float)\n        data = TensorDataset(all_source_ids, all_target_ids)\n    return examples, data ,df\n\n\ndef read_examples(filename, data_num, task):\n    read_example_dict = {\n        'sort' : read_sort_examples,\n    }\n    return read_example_dict[task](filename, data_num)","metadata":{"_uuid":"ff731a57-591e-41b7-b5c7-270663678069","_cell_guid":"23fe62f3-e35e-4011-b883-0dc09c97dd6a","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:47:17.831705Z","iopub.execute_input":"2022-07-21T13:47:17.832321Z","iopub.status.idle":"2022-07-21T13:47:17.844311Z","shell.execute_reply.started":"2022-07-21T13:47:17.832287Z","shell.execute_reply":"2022-07-21T13:47:17.843541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def clean_code(cell):\n    return str(cell).replace(\"\\\\n\", \"\\n\")\n\n\ndef sample_cells(cells, n):\n    cells = [clean_code(cell) for cell in cells]\n    if n >= len(cells):\n        return [cell[:200] for cell in cells]\n    else:\n        results = []\n        step = len(cells) / n\n        idx = 0\n        while int(np.round(idx)) < len(cells):\n            results.append(cells[int(np.round(idx))])\n            idx += step\n        assert cells[0] in results\n        if cells[-1] not in results:\n            results[-1] = cells[-1]\n        return results\n\n\ndef get_features(df):\n    features = dict()\n    df = df.sort_values(\"rank\").reset_index(drop=True)\n    for idx, sub_df in tqdm(df.groupby(\"id\")):\n        features[idx] = dict()\n        total_md = sub_df[sub_df.cell_type == \"markdown\"].shape[0]\n        code_sub_df = sub_df[sub_df.cell_type == \"code\"]\n        total_code = code_sub_df.shape[0]\n        codes = sample_cells(code_sub_df.source.values, 50)\n        features[idx][\"total_code\"] = total_code\n        features[idx][\"total_md\"] = total_md\n        features[idx][\"codes\"] = codes\n    return features","metadata":{"_uuid":"a9979e54-02f2-4e89-875a-f7636da9fbcc","_cell_guid":"f2c2bb63-8cde-423c-b4bb-eafad698c6a6","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:47:18.026179Z","iopub.execute_input":"2022-07-21T13:47:18.026526Z","iopub.status.idle":"2022-07-21T13:47:18.037575Z","shell.execute_reply.started":"2022-07-21T13:47:18.026497Z","shell.execute_reply":"2022-07-21T13:47:18.036553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_sort_examples(data_path, num):\n    \"\"\"Read examples from filename.\"\"\"\n    examples = []\n    train_df_mark=pd.read_csv(data_path).drop(\"parent_id\", axis=1).dropna().reset_index(drop=True)        \n    for i in range(num):\n            x = train_df_mark.iloc[i]\n            examples.append(\n                Example(\n                    idx=x[\"id\"],\n                    source=x[\"source\"].strip(),\n                    target=x[\"pct_rank\"]\n                )\n            )\n    return examples,train_df_mark","metadata":{"_uuid":"7a6c1185-55a1-44f8-aea8-7b2258e871e9","_cell_guid":"f8d60e6d-c384-4c43-ac12-982c0e8c9ab2","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:47:18.249996Z","iopub.execute_input":"2022-07-21T13:47:18.250613Z","iopub.status.idle":"2022-07-21T13:47:18.258092Z","shell.execute_reply.started":"2022-07-21T13:47:18.250580Z","shell.execute_reply":"2022-07-21T13:47:18.255516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def convert_examples_to_features(item):\n    stage = ''\n    example, example_index,tokenizer, fts,total_max_len = item\n\n    if 'codet5' in ['t5', 'codet5'] and True :\n            source_str = \"{}: {}\".format('sort', example.source)\n    else:\n        source_str = example.source\n        \n    source_str = source_str.replace('</s>', '<unk>')\n    source_ids = tokenizer.encode(source_str, max_length=512, padding='max_length', truncation=True)\n    if source_ids.count(tokenizer.eos_token_id) == 1:\n        target_ids = example.target\n        code_inputs = tokenizer.batch_encode_plus(\n                [str(x) for x in fts[example.id][\"codes\"]],\n                add_special_tokens=True,\n                max_length=512,\n                padding=\"max_length\",\n                truncation=True\n            )\n        n_md = fts[example.id][\"total_md\"]\n        n_code = fts[example.id][\"total_md\"]\n        if n_md + n_code == 0:\n            fts = torch.FloatTensor([0])\n        else:\n            fts = torch.FloatTensor([n_md / (n_md + n_code)])\n        ids = source_ids\n\n        for x in code_inputs:\n            ids.extend(x[:-1])\n        ids = ids[:total_max_len]\n        if len(ids) != total_max_len:\n            ids = ids + [tokenizer.pad_token_id, ] * (total_max_len - len(ids))\n    \n        assert len(ids) == total_max_len\n    \n\n        return InputFeatures(\n            example_index,\n            ids,\n            target_ids,\n            fts\n        )\n    else : \n        return None","metadata":{"_uuid":"448a9412-28f5-42d2-858e-fd7b310b0c98","_cell_guid":"da8e5b9b-c437-444b-b5ec-2835022a2b7f","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:47:20.027719Z","iopub.execute_input":"2022-07-21T13:47:20.028151Z","iopub.status.idle":"2022-07-21T13:47:20.039210Z","shell.execute_reply.started":"2022-07-21T13:47:20.028118Z","shell.execute_reply":"2022-07-21T13:47:20.038359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_ranks(base, derived):\n    return [base.index(d) for d in derived]","metadata":{"_uuid":"facedb1a-94b2-41c0-aa87-fe4de38fec2c","_cell_guid":"1ff16bff-e02c-41c6-a49b-8976d3d85043","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:47:20.909409Z","iopub.execute_input":"2022-07-21T13:47:20.910083Z","iopub.status.idle":"2022-07-21T13:47:20.914060Z","shell.execute_reply.started":"2022-07-21T13:47:20.910051Z","shell.execute_reply":"2022-07-21T13:47:20.913202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"    import multiprocessing\n    from torch.utils.tensorboard import SummaryWriter\n    import pandas as pd\n    from pathlib import Path\n    import time\n    #parser = argparse.ArgumentParser()\n    #args = add_args(parser)\n    #logger.info(args)\n    t0 = time.time()\n    data_dir = Path('../input/AI4Code')\n    NUM_TRAIN = 5500\n\n\n    def read_notebook(path):\n        return (\n            pd.read_json(\n                path,\n                dtype={'cell_type': 'category', 'source': 'str'})\n            .assign(id=path.stem)\n            .rename_axis('cell_id')\n        )\n\n    \n    paths_train = list((data_dir / 'train').glob('*.json'))[4000:NUM_TRAIN]\n    notebooks_train = [\n        read_notebook(path) for path in tqdm(paths_train, desc='Train NBs')\n    ]\n    df = (\n        pd.concat(notebooks_train)\n        .set_index('id', append=True)\n        .swaplevel()\n        .sort_index(level='id', sort_remaining=False)\n    )\n    df_orders = pd.read_csv(\n    data_dir / 'train_orders.csv',\n    index_col='id',\n    squeeze=True,\n    ).str.split()  # Split the string representation of cell_ids into a list\n\n\n    df_orders_ = df_orders.to_frame().join(\n    df.reset_index('cell_id').groupby('id')['cell_id'].apply(list),\n    how='right',\n    )\n\n    ranks = {}\n    for id_, cell_order, cell_id in df_orders_.itertuples():\n        ranks[id_] = {'cell_id': cell_id, 'rank': get_ranks(cell_order, cell_id)}\n\n    df_ranks = (\n        pd.DataFrame\n        .from_dict(ranks, orient='index')\n        .rename_axis('id')\n        .apply(pd.Series.explode)\n        .set_index('cell_id', append=True)\n    )","metadata":{"_uuid":"6be55a0f-a119-4701-aea6-9d90e0c79c96","_cell_guid":"7579ea5f-d655-4c0a-975e-eadcbec335d7","collapsed":false,"jupyter":{"outputs_hidden":false},"scrolled":true,"execution":{"iopub.status.busy":"2022-07-21T13:47:50.131186Z","iopub.execute_input":"2022-07-21T13:47:50.132090Z","iopub.status.idle":"2022-07-21T13:48:08.691708Z","shell.execute_reply.started":"2022-07-21T13:47:50.132043Z","shell.execute_reply":"2022-07-21T13:48:08.690905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_ancestors = pd.read_csv(data_dir / 'train_ancestors.csv', index_col='id')\ndf_ancestors","metadata":{"_uuid":"00f4fa78-f08f-4eeb-a064-1def8470d707","_cell_guid":"0dc7b1f3-0a11-4b38-8d38-67b2b9d02ab7","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:48:08.693352Z","iopub.execute_input":"2022-07-21T13:48:08.693783Z","iopub.status.idle":"2022-07-21T13:48:08.906648Z","shell.execute_reply.started":"2022-07-21T13:48:08.693748Z","shell.execute_reply":"2022-07-21T13:48:08.905848Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = df.reset_index().merge(df_ranks, on=[\"id\", \"cell_id\"]).merge(df_ancestors, on=[\"id\"])\ndf","metadata":{"_uuid":"6555d900-1cdd-46ba-8ae0-dd6b5e317ad6","_cell_guid":"f69b0dd7-a15d-4669-a19f-25b4e9866d01","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:48:08.907807Z","iopub.execute_input":"2022-07-21T13:48:08.908601Z","iopub.status.idle":"2022-07-21T13:48:09.043018Z","shell.execute_reply.started":"2022-07-21T13:48:08.908557Z","shell.execute_reply":"2022-07-21T13:48:09.042112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[\"pct_rank\"] = df[\"rank\"] / df.groupby(\"id\")[\"cell_id\"].transform(\"count\")\ndf","metadata":{"_uuid":"94aaa580-f4a7-414e-8bb3-627ae02fb5f9","_cell_guid":"bdf72275-9642-48fd-8e6d-f8689a7923b5","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:48:09.045165Z","iopub.execute_input":"2022-07-21T13:48:09.045529Z","iopub.status.idle":"2022-07-21T13:48:09.086536Z","shell.execute_reply.started":"2022-07-21T13:48:09.045495Z","shell.execute_reply":"2022-07-21T13:48:09.085788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AdamW, get_linear_schedule_with_warmup\n\n\nconfig_class, model_class, tokenizer_class = build_or_load_gen_model()\nconfig = config_class.from_pretrained(\"Salesforce/codet5-small\")\nmodel = model_class.from_pretrained(\"Salesforce/codet5-small\")\ntokenizer = tokenizer_class.from_pretrained(\"Salesforce/codet5-small\")\n    \nmodel = SortModel(model, config, tokenizer)\ntorch.device('cuda',0)\n        #model.to(device)\n        #pool = multiprocessing.Pool(multiprocessing.cpu_count())\n        #train_filename, dev_filename, test_filename = get_filenames('../input/dataai4code', args.task, args.sub_task)\n        #fa = open(os.path.join('./outputs/', 'summary.log'), 'a+')\n\n        #if args.do_train:\n            #if args.local_rank in [-1, 0] and args.data_num == -1:\n\nmodel = nn.DataParallel(model)\nmodel.load_state_dict(torch.load('../input/m20000-1/model_20000_1.bin'))\nsummary_fn = './outputs/summary'\ntb_writer = SummaryWriter(summary_fn)\nfrom sklearn.model_selection import GroupShuffleSplit\n\nNVALID = 0.1  # size of validation set\n\nsplitter = GroupShuffleSplit(n_splits=1, test_size=NVALID, random_state=0)\n    #train_examples, train_data,df = load_and_cache_gen_data(3000,'../input/dataai4code/data/train_mark.csv', tokenizer, ' ', only_src=False, is_sample=False)\ntrain_ind, val_ind = next(splitter.split(df, groups=df[\"ancestor_id\"]))\ntrain_df = df.loc[train_ind].reset_index(drop=True)\nval_df = df.loc[val_ind].reset_index(drop=True)\n    \ntrain_df_mark = train_df[train_df[\"cell_type\"] == \"markdown\"].reset_index(drop=True)\nval_df_mark = val_df[val_df[\"cell_type\"] == \"markdown\"].reset_index(drop=True)\n    # Prepare training data loader","metadata":{"_uuid":"83d1f7c1-2b83-4106-96c7-dd7ae5e68952","_cell_guid":"2a5f5dcd-7877-40e0-9667-2489d1ea2d98","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T14:02:16.705756Z","iopub.execute_input":"2022-07-21T14:02:16.706472Z","iopub.status.idle":"2022-07-21T14:02:26.644515Z","shell.execute_reply.started":"2022-07-21T14:02:16.706430Z","shell.execute_reply":"2022-07-21T14:02:26.643678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fts=get_features(df)","metadata":{"_uuid":"bddbbf58-5312-4a4b-849a-592d94d7ec8e","_cell_guid":"9b78e4d8-e72b-44cb-ad4e-f3213041ebe8","collapsed":false,"scrolled":true,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:48:53.122721Z","iopub.execute_input":"2022-07-21T13:48:53.123079Z","iopub.status.idle":"2022-07-21T13:48:54.575241Z","shell.execute_reply.started":"2022-07-21T13:48:53.123044Z","shell.execute_reply":"2022-07-21T13:48:54.574333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"examples = []\nfor i in range(len(train_df_mark)):\n            x = train_df_mark.iloc[i]\n            examples.append(\n                Example(\n                    idx=x[\"id\"],\n                    source=x[\"source\"].strip(),\n                    target=x[\"pct_rank\"]\n                )\n            )\ntuple_examples = [(example, idx, tokenizer,fts,512) for idx, example in enumerate(examples)]\nfeatures = list(map(convert_examples_to_features, tqdm(tuple_examples, total=len(tuple_examples))))\na= features[0]\nall_source_ids = torch.tensor([f.source_ids for f in features], dtype=torch.long)\n#all_fts = torch.tensor([f.fts for f in features], dtype=torch.long)\nall_target_ids = torch.tensor([f.target_ids for f in features], dtype=torch.float)\ndata = TensorDataset(all_source_ids, all_target_ids.unsqueeze(dim=1))","metadata":{"_uuid":"e7b9d7ee-996f-42b4-94c0-5c6673ca0959","_cell_guid":"2600f887-c1f9-49d3-b5ba-7afb850a2fef","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:48:54.577119Z","iopub.execute_input":"2022-07-21T13:48:54.577748Z","iopub.status.idle":"2022-07-21T13:59:04.452276Z","shell.execute_reply.started":"2022-07-21T13:48:54.577709Z","shell.execute_reply":"2022-07-21T13:59:04.451492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data","metadata":{"_uuid":"26adde8e-3d2e-4ea9-9abb-cbe018c67077","_cell_guid":"5a3c8369-9f34-4fbb-a904-85e430fdb5e3","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-19T14:11:46.178655Z","iopub.execute_input":"2022-07-19T14:11:46.180885Z","iopub.status.idle":"2022-07-19T14:11:46.189890Z","shell.execute_reply.started":"2022-07-19T14:11:46.180844Z","shell.execute_reply":"2022-07-19T14:11:46.189167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"examples = []\nfor i in range(len(val_df_mark)):\n            x = val_df_mark.iloc[i]\n            examples.append(\n                Example(\n                    idx=x[\"id\"],\n                    source=x[\"source\"].strip(),\n                    target=x[\"pct_rank\"]\n                )\n            )\ntuple_examples = [(example, idx, tokenizer,fts,512) for idx, example in enumerate(examples)]\nfeatures = list(map(convert_examples_to_features, tqdm(tuple_examples, total=len(tuple_examples))))\nall_source_ids = torch.tensor([f.source_ids for f in features], dtype=torch.long)\n#all_fts = torch.tensor([f.fts for f in features], dtype=torch.long)\nall_target_ids = torch.tensor([f.target_ids for f in features], dtype=torch.float)\ndata_val = TensorDataset(all_source_ids, all_target_ids.unsqueeze(dim=1))","metadata":{"_uuid":"c2c37906-9c52-4b67-a7b2-b2e780ec0b0e","_cell_guid":"96489da5-951f-4dc7-b82e-a6a5fcc3fea0","collapsed":false,"jupyter":{"outputs_hidden":false},"scrolled":true,"execution":{"iopub.status.busy":"2022-07-21T13:59:04.454596Z","iopub.execute_input":"2022-07-21T13:59:04.455434Z","iopub.status.idle":"2022-07-21T13:59:50.111351Z","shell.execute_reply.started":"2022-07-21T13:59:04.455394Z","shell.execute_reply":"2022-07-21T13:59:50.110503Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_dataloader = DataLoader(data_val, batch_size=4,num_workers=2, pin_memory=True)","metadata":{"_uuid":"08b9cf10-b11b-400d-b876-8fb72bf14585","_cell_guid":"8e4ff789-6ccc-4cec-becb-66c556ea20dc","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:59:50.112890Z","iopub.execute_input":"2022-07-21T13:59:50.113512Z","iopub.status.idle":"2022-07-21T13:59:50.118470Z","shell.execute_reply.started":"2022-07-21T13:59:50.113475Z","shell.execute_reply":"2022-07-21T13:59:50.117612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_dataloader","metadata":{"_uuid":"965da267-b1d0-4506-9caa-3414e6a39707","_cell_guid":"01f695bc-14ec-434f-b5c6-cdba80b399de","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:59:50.119733Z","iopub.execute_input":"2022-07-21T13:59:50.121201Z","iopub.status.idle":"2022-07-21T13:59:50.132179Z","shell.execute_reply.started":"2022-07-21T13:59:50.121165Z","shell.execute_reply":"2022-07-21T13:59:50.131145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def adjust_lr(optimizer, epoch):\n    if epoch < 1:\n        lr = 5e-5\n    elif epoch < 2:\n        lr = 5e-5\n    elif epoch < 5:\n        lr = 5e-5\n    else:\n        lr = 1e-5\n\n    for p in optimizer.param_groups:\n        p['lr'] = lr\n    return lr\n    \ndef get_optimizer(net):\n    optimizer = torch.optim.Adam(filter(lambda p: p.requires_grad, net.parameters()), lr=3e-4, betas=(0.9, 0.999),\n                                 eps=1e-8 )\n    return optimizer","metadata":{"_uuid":"8b615150-4dbd-4126-a503-42adb08040ce","_cell_guid":"3d5ac4a2-dc68-46c8-8f7d-2835c912ecbb","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T13:59:50.133687Z","iopub.execute_input":"2022-07-21T13:59:50.134169Z","iopub.status.idle":"2022-07-21T13:59:50.143173Z","shell.execute_reply.started":"2022-07-21T13:59:50.134131Z","shell.execute_reply":"2022-07-21T13:59:50.142191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"    import gc\n    import matplotlib.pyplot as plt\n    import os\n\n    #train_sampler = RandomSampler(train_data) if args.local_rank == -1 else DistributedSampler(train_data)\n    train_dataloader = DataLoader(data, batch_size=4,\n                                      num_workers=2, pin_memory=True)\n    # Prepare optimizer and schedule (linear warmup and decay)\n    no_decay = ['bias', 'LayerNorm.weight']\n    optimizer_grouped_parameters = [\n             {'params': [p for n, p in model.named_parameters() if not any(nd in n for nd in no_decay)],\n              'weight_decay': 0},\n             {'params': [p for n, p in model.named_parameters() if any(nd in n for nd in no_decay)], 'weight_decay': 0.0}\n         ]\n\n    num_train_optimization_steps = 5 * len(train_dataloader)\n    optimizer = get_optimizer(model)\n    #optimizer = AdamW(optimizer_grouped_parameters, lr=5e-5, eps=1e-8)\n    scheduler = get_linear_schedule_with_warmup(optimizer, num_warmup_steps=100, num_training_steps=num_train_optimization_steps)\n    \n    gc.collect()\n    torch.cuda.empty_cache()\n\n    os.environ['CUDA_LAUNCH_BLOCKING'] = '1'\n    device = torch.device('cuda', 0)\n    model.to(device)\n    dev_dataset = {}\n    # running_loss= 0.0\n    loss_values = []\n    #global_step, best_bleu_em, best_ppl = 0, -1, 1e6\n    #not_loss_dec_cnt, not_bleu_em_inc_cnt = 0, 1e6\n    \n    res_per_epoch = dict()\n    \n    for cur_epoch in range(3):\n        model.train()\n        \n        bar = tqdm(train_dataloader, total=len(train_dataloader), desc=\"Training\")\n        # nb_tr_examples, nb_tr_steps, tr_loss = 0, 0, 0\n        lr = adjust_lr(optimizer, cur_epoch)\n        \n        train_loss = 0.0\n        \n        for step, data in enumerate(bar):\n\n            #print(f\"Step n°{step}\\n\")\n            #print(f\"DATA: \\n{data}\\n\")\n            data = tuple(t.to(device) for t in data)\n            source_ids, labels = data # inputs tokénisés et pct_rank véridiques \n            \n            # Clear the gradients\n            \n            accumulation_steps = 1.2   \n            \n            # Forward pass and find the loss\n            loss, prob = model(input_ids=source_ids, labels=labels,fts=fts)\n            # nb_tr_examples += source_ids.size(0)\n            # nb_tr_steps += 1\n            loss.backward()\n\n            \n            # Calculate gradients\n            # Update weights\n            optimizer.step()\n            optimizer.zero_grad()\n            scheduler.step()\n\n            train_loss += loss.item()\n            #print(f\"\\n\\033[01mInputs\\033[0m:\\n{source_ids}\\n\\n\\033[01mTarget\\033[0m:\\n{labels}\\n\\n\\033[01mPrédictions\\033[0m:\\n{prob[:,1]}\\n\")\n            bar.set_description(f\"Epoch {cur_epoch+1} Loss: {train_loss / len(train_dataloader)} lr: {lr}\")\n        #wandb.log({\"loss\": train_loss})\n\n        train_loss = train_loss / len(train_dataloader)\n        \n        print(f\"Training loss: {train_loss}\\n\")\n        \n        loss_values.append(train_loss)","metadata":{"_uuid":"d8dc175e-d2df-4333-9bfe-a4792c84eb8d","_cell_guid":"c27648df-eece-4498-99fd-e849d826b20a","collapsed":false,"scrolled":true,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T14:02:26.646165Z","iopub.execute_input":"2022-07-21T14:02:26.646515Z","iopub.status.idle":"2022-07-21T14:02:37.322980Z","shell.execute_reply.started":"2022-07-21T14:02:26.646488Z","shell.execute_reply":"2022-07-21T14:02:37.321341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(loss_values)","metadata":{"_uuid":"c5ad4864-caf8-4f85-8c67-3e3e92ca4cd1","_cell_guid":"898330bc-7b7e-49c1-b514-9134928dd516","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T14:00:00.024216Z","iopub.status.idle":"2022-07-21T14:00:00.024900Z","shell.execute_reply.started":"2022-07-21T14:00:00.024644Z","shell.execute_reply":"2022-07-21T14:00:00.024673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(loss_values)","metadata":{"_uuid":"5dbfc8b8-68ef-408f-bd5d-0c96893a08c5","_cell_guid":"67b04b6b-8379-439f-a73a-486f4f7e5648","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T14:00:00.026464Z","iopub.status.idle":"2022-07-21T14:00:00.026910Z","shell.execute_reply.started":"2022-07-21T14:00:00.026686Z","shell.execute_reply":"2022-07-21T14:00:00.026709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"_uuid":"0198a7f0-7ad8-4743-8ba7-8f6731e39854","_cell_guid":"6a47ec09-e684-49ec-b1bf-1ac6477c0906","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\ndevice = torch.device('cuda', 0)\nmodel.to(device)","metadata":{"_uuid":"81aaab1f-a835-40f6-b632-d88932ca2a80","_cell_guid":"2140fc6a-45f9-479f-8a5b-6d9b3bf56cdf","collapsed":false,"scrolled":true,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T14:00:36.866348Z","iopub.execute_input":"2022-07-21T14:00:36.866760Z","iopub.status.idle":"2022-07-21T14:00:36.887133Z","shell.execute_reply.started":"2022-07-21T14:00:36.866716Z","shell.execute_reply":"2022-07-21T14:00:36.886245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":" torch.save(model.state_dict(), \"./outputs/model2.bin\")","metadata":{"_uuid":"42f858e0-6561-4de9-a1b6-fc5da0c38836","_cell_guid":"08622673-1366-4004-a9e0-02ccad0d0da1","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#_,data,val = load_and_cache_gen_data(10000,'../input/dataai4code/data/val_mark.csv',tokenizer,'')","metadata":{"_uuid":"61d0d196-555d-4af3-a010-3fcee00019ff","_cell_guid":"a6e8ad07-a761-4a5d-b857-594facf0ef55","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-07T15:29:41.680999Z","iopub.execute_input":"2022-07-07T15:29:41.681397Z","iopub.status.idle":"2022-07-07T15:29:41.685729Z","shell.execute_reply.started":"2022-07-07T15:29:41.681359Z","shell.execute_reply":"2022-07-07T15:29:41.684659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#_,data2,val2 = load_and_cache_gen_data(200,'../input/dataai4code/data/val.csv',tokenizer,'')","metadata":{"_uuid":"5ff75203-e942-4e68-ae01-2433e20b7de8","_cell_guid":"7194b4b0-16d9-4899-9918-3163a647f0cb","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-07T15:29:41.687047Z","iopub.execute_input":"2022-07-07T15:29:41.687411Z","iopub.status.idle":"2022-07-07T15:29:41.696582Z","shell.execute_reply.started":"2022-07-07T15:29:41.687376Z","shell.execute_reply":"2022-07-07T15:29:41.695869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#val_code = val2[val2['cell_type']=='code']\n#val_code","metadata":{"_uuid":"3ee3031b-4c6f-43e2-ae56-5069750fdf74","_cell_guid":"004e4dde-3283-458a-95ed-8a3fa42b3cb3","collapsed":false,"execution":{"iopub.status.busy":"2022-07-06T13:14:29.18174Z","iopub.execute_input":"2022-07-06T13:14:29.182363Z","iopub.status.idle":"2022-07-06T13:14:29.191632Z","shell.execute_reply.started":"2022-07-06T13:14:29.182321Z","shell.execute_reply":"2022-07-06T13:14:29.19056Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#val_mark = val2[val2['cell_type']=='markdown']\n#val_mark","metadata":{"_uuid":"dfa0e754-7abf-4e8b-8813-f44a9d0ed11e","_cell_guid":"490a0946-de5f-4499-9b9d-b04e483044cb","collapsed":false,"execution":{"iopub.status.busy":"2022-07-06T13:14:29.193048Z","iopub.execute_input":"2022-07-06T13:14:29.194173Z","iopub.status.idle":"2022-07-06T13:14:29.20257Z","shell.execute_reply.started":"2022-07-06T13:14:29.194123Z","shell.execute_reply":"2022-07-06T13:14:29.200931Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#data = torch.load('../input/tensors/tensor(1).pt')\n#len(data.tensors[0])","metadata":{"_uuid":"1a1753cd-86a6-44d0-828f-7e614d7428ae","_cell_guid":"aecc9484-a1be-4f58-9b52-8cfb24b753f9","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''config_class, model_class, tokenizer_class = build_or_load_gen_model()\nconfig = config_class.from_pretrained(\"Salesforce/codet5-small\")\nmodel = model_class.from_pretrained(\"Salesforce/codet5-small\")\ntokenizer = tokenizer_class.from_pretrained(\"Salesforce/codet5-small\")\n    \n    \n\nmodel = SortModel(model, config, tokenizer)\n#model = nn.DataParallel(model)\ndevice = torch.device('cuda', 0)\nmodel.to(device)\n\nmodel.load_state_dict(torch.load('../input/m7000-10/model(2).bin'))'''","metadata":{"_uuid":"98d1f390-8ba7-46f5-8b56-aee7a1794691","_cell_guid":"619b4ebd-1d62-44c9-8cc3-f4353c65b3dc","collapsed":false,"jupyter":{"outputs_hidden":false},"scrolled":true,"execution":{"iopub.status.busy":"2022-07-19T09:43:58.293658Z","iopub.execute_input":"2022-07-19T09:43:58.294112Z","iopub.status.idle":"2022-07-19T09:44:05.589242Z","shell.execute_reply.started":"2022-07-19T09:43:58.294063Z","shell.execute_reply":"2022-07-19T09:44:05.588440Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Eval!\nlogger.info(\"***** Running evaluation *****\")\n#logger.info(\"  Num examples = %d\", len(eval_examples))\nlogger.info(\"  Num batches = %d\", len(val_dataloader))\n#logger.info(\"  Batch size = %d\", args.eval_batch_size)\neval_loss = 0.0\nnb_eval_steps = 0\nmodel.eval()\n\nlogits = []\nlabels = []\nfor batch in tqdm(val_dataloader, total=len(val_dataloader), desc=\"Evaluating\"):\n        inputs = batch[0].to(device)\n        label = batch[1].to(device)\n        with torch.no_grad():\n            lm_loss, logit = model(inputs, label)\n            eval_loss += lm_loss.mean().item()\n            logits.append(logit.cpu().numpy())\n            labels.append(label.cpu().numpy())\n        nb_eval_steps += 1\nlogits = np.concatenate(logits, 0)\nlabels = np.concatenate(labels, 0)","metadata":{"_uuid":"fefbeb42-a1b7-46e9-ab06-6ee161a48007","_cell_guid":"795268ca-ec38-43b9-a349-caa5f573c9cc","collapsed":false,"scrolled":true,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T14:00:00.028277Z","iopub.status.idle":"2022-07-21T14:00:00.028950Z","shell.execute_reply.started":"2022-07-21T14:00:00.028647Z","shell.execute_reply":"2022-07-21T14:00:00.028688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(logits,'./outputs/logits.pt')","metadata":{"_uuid":"dc371b58-b993-4cef-a179-4457a2582974","_cell_guid":"b223b696-a6e8-4233-a63f-902f3ac0eb03","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T09:15:24.009586Z","iopub.execute_input":"2022-07-21T09:15:24.010129Z","iopub.status.idle":"2022-07-21T09:15:24.016301Z","shell.execute_reply.started":"2022-07-21T09:15:24.010090Z","shell.execute_reply":"2022-07-21T09:15:24.015523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(labels,'./outputs/labels.pt')","metadata":{"_uuid":"27e4e51a-d6b1-414c-a822-e188b68a01bd","_cell_guid":"d9d3fc63-6e26-4331-9ac6-723d8465dd78","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T09:15:24.017736Z","iopub.execute_input":"2022-07-21T09:15:24.018154Z","iopub.status.idle":"2022-07-21T09:15:24.028179Z","shell.execute_reply.started":"2022-07-21T09:15:24.018118Z","shell.execute_reply":"2022-07-21T09:15:24.027168Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#logits = torch.load(\"../input/logits/logits.pt\")","metadata":{"_uuid":"1bbef771-4f76-4600-a12d-62453abff0d2","_cell_guid":"decd881b-4d47-49c5-ae1f-4695f549278e","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T09:15:24.029627Z","iopub.execute_input":"2022-07-21T09:15:24.030113Z","iopub.status.idle":"2022-07-21T09:15:24.037362Z","shell.execute_reply.started":"2022-07-21T09:15:24.030076Z","shell.execute_reply":"2022-07-21T09:15:24.036385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#labels = torch.load(\"../input/labels/labels.pt\")","metadata":{"_uuid":"8d525bbb-6e01-4124-b87a-eadb897cf374","_cell_guid":"f6f4b387-a6d4-40e8-88c2-934d95705acb","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T09:15:24.040114Z","iopub.execute_input":"2022-07-21T09:15:24.040461Z","iopub.status.idle":"2022-07-21T09:15:24.046398Z","shell.execute_reply.started":"2022-07-21T09:15:24.040425Z","shell.execute_reply":"2022-07-21T09:15:24.045653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_df[\"pred\"] = val_df.groupby([\"id\", \"cell_type\"])[\"rank\"].rank(pct=True)\nval_df.loc[val_df[\"cell_type\"] == \"markdown\", \"pred\"] = logits","metadata":{"_uuid":"32abd748-9ed6-4ce9-a033-c80091474d51","_cell_guid":"e75fa438-a2e6-4dd7-9f63-24182e1b90ae","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T14:00:00.030556Z","iopub.status.idle":"2022-07-21T14:00:00.031344Z","shell.execute_reply.started":"2022-07-21T14:00:00.030981Z","shell.execute_reply":"2022-07-21T14:00:00.031004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = val_df.sort_values(\"pred\").groupby(\"id\")[\"cell_id\"].apply(lambda x: \" \".join(x)).reset_index()\nsub_df.rename(columns={\"cell_id\": \"cell_id\"}, inplace=True)\nsub_df.head()","metadata":{"_uuid":"bcd5c952-435f-497e-91fe-c87d3f3dcd2d","_cell_guid":"c7dd9241-a25e-4c1a-9032-d11c0f76c03c","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T14:00:00.032861Z","iopub.status.idle":"2022-07-21T14:00:00.033331Z","shell.execute_reply.started":"2022-07-21T14:00:00.033107Z","shell.execute_reply":"2022-07-21T14:00:00.033129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Metric\nfrom bisect import bisect\n\ndef count_inversions(a):\n    inversions = 0\n    sorted_so_far = []\n    for i, u in enumerate(a):\n        j = bisect(sorted_so_far, u)\n        inversions += i - j\n        sorted_so_far.insert(j, u)\n    return inversions\n\n\ndef kendall_tau(ground_truth, predictions):\n    total_inversions = 0\n    total_2max = 0  # twice the maximum possible inversions across all instances\n    for gt, pred in zip(ground_truth, predictions):\n        ranks = [gt.index(x) for x in pred]  # rank predicted order in terms of ground truth\n        total_inversions += count_inversions(ranks)\n        n = len(gt)\n        total_2max += n * (n - 1)\n    return 1 - 4 * total_inversions / total_2max","metadata":{"_uuid":"d4a1869d-9446-46d5-9908-18e6c058cb5b","_cell_guid":"cd86af28-241b-41a4-8cf9-b6d4b276ca3c","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T14:00:00.034446Z","iopub.status.idle":"2022-07-21T14:00:00.034970Z","shell.execute_reply.started":"2022-07-21T14:00:00.034703Z","shell.execute_reply":"2022-07-21T14:00:00.034729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred= val_df.sort_values(\"pred\").groupby(\"id\")[\"cell_id\"].apply(list)\ngd = df_orders.loc[pred.index]\nacc = kendall_tau(gd,pred)\nacc","metadata":{"_uuid":"6aa550bd-3efe-43c8-a987-978f7dfd82fa","_cell_guid":"f5315a35-14b6-4f00-b87f-5bcc6573f6ef","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-21T14:00:00.038655Z","iopub.status.idle":"2022-07-21T14:00:00.039201Z","shell.execute_reply.started":"2022-07-21T14:00:00.038908Z","shell.execute_reply":"2022-07-21T14:00:00.038933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(acc,'./outputs/acca.pt')","metadata":{"_uuid":"ae6a16f6-f861-435e-8bd6-bc3e350ff500","_cell_guid":"134e48de-1c7f-421f-9fd4-1951a0d95d98","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]}]}