{"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":"# FULL CREDIT TO THE BELOW AUTHORS\n# CHRIS DEOTTE\n# Y. NAKAMA\n# CPMP","metadata":{"execution":{"iopub.status.busy":"2023-06-23T20:08:56.344857Z","iopub.execute_input":"2023-06-23T20:08:56.345192Z","iopub.status.idle":"2023-06-23T20:08:56.352968Z","shell.execute_reply.started":"2023-06-23T20:08:56.34516Z","shell.execute_reply":"2023-06-23T20:08:56.352032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Protein BERT Starter Notebook - LB 0.20+\n\n## Torch DataParallel Use\n\nThis notebooks shows how to use torch DistibutedParallel (DP) to use two T4 at once. It is an adaptation of Chris Deotte notebook you can find at https://www.kaggle.com/code/cdeotte/protein-bert-finetune-lb-0-30. All content is from Chris or the people he forked from except for DP related code. The main changes are:\n\n* All model inputs are passed as a single dictionary. DP works fine if the model input is a dictionary of tensors, or if the model inputs are several tensors. DP data split between GPU breaks if you mix the two ways of passing input. The code was changed in data loader and training/evalution loops.\n* Batch size is doubled. It may not be the best for model accuracy and other parameters may need to be tuned again. \n* Data loaders us 2 workeds instead of zero.\n* Parallelism is disabled for tokenizers. Issues with tokenizer parallelism is the primary reason people use zero workers in data loaders.\n\nI used comments containing CPMP to indicate my code changes.\n\nDP works fine except if you want to also use gradient checkpointing. Using gradient checkpointing with DP is extremely slow. It is better to use DistributedDataParallel (DDP) for gradient checkpointing. However DDP cannot be used in a notebook. I will work on a DDP script version of this notebook later on.\n\n\n## Original Introduction\n\nThis notebook demonstrates how to finetune an NLP transformer to compete in Kaggle's Novozymes Enzyme Stability Prediction competition. We train with Jin's external data described [here][2] which uses delta `Tm` and delta `dG` targets. How to train with `dTm` target instead of `Tm` target is explained [here][4]. We use the PyTorch pipeline from Y.Nakama's starter NLP notebook from Feedback Prize 3 competition [here][1]. We will finetune HuggingFace's pretrained `Rostlab/prot_bert` model.\n\nThe competition FP3 is a regression NLP task. Our competition NESP is also a regression task and we can treat protein amino acid sequences as \"sentences\" where each amino acid is a \"word\". \n\nThe secret ingredient for success is the architecture of our model. In FP3 we feed one sentence and get one regression. In NESP, we must feed two \"sentences\". We input both the wild type sequence and mutant sequence into our model. Then we subtract the embeddings and concatenate that result with the wild type and mutant embeddings. We do this both with the specific mutant amino acid token output embedding, and the entire sequences embedding after mean pooling. Finally we use a dense layer to predict one regression target. See diagram below and review code in code cell 23.\n\n![](https://raw.githubusercontent.com/cdeotte/Kaggle_Images/main/Oct-2022/prot_bert.png)\n\nFurthermore, we freeze 22 out of the 30 layers of Protein BERT to retain most of Protein BERT original pretraining knowledge which improves CV LB. By modifying the architecture, freezing, training schedule, and other hyperparameters, it is possible to improve this notebook's CV LB. Also we can use different train data like FireProtDB [here][6]. Also we can try using different pretrained transformers besides `Rostlab/prot_bert` like Facebooks's ESM or MSA [here][5]. This notebook trains 3 out of 5 folds. Using more folds improves CV LB.\n\nCurrently we are using Kaggle's new 1xT4 GPU instead of 1xP100 GPU. This gives us 1.7x speed up. Probably because of using APEX mixed precision. If someone can get this notebook to use both of Kaggle's new 2xT4 GPUs that would be great! We would get another 2x speed and 2x memory boost. I tried using `torch.nn.DataParallel` but i get some error messages. Note that using `gradient_checkpointing=True` (with `prot_bert`) produces bad CV LB. I'm not sure why. \n\nThe processed training data for this notebook comes from my XGB notebook [here][3]. In code cell 9, we save the dataframe after `df = pd.concat([df,df2,df3,kaggle])`\n\n## Notes:\n\n* **Version 1-3** finetunes HuggingFace's `Rostlab/prot_bert` 3 out of 5 folds and achieves single model LB 0.153. If we train with `batch_size=16` and `LR=1e-5` then single fold 0 achieves LB 0.209 (confirmed with offline training). However in Kaggle notebook, i cannot train with that large batch size and `gradient_checkpointing` nor `gradient_accumulation_steps` replicates large batch results. Not sure why. \n* **Version 4** finetunes HuggingFace's `facebook/esm2_t33_650M_UR50D` 5 out of 5 folds and achieves LB ???. We freeze 32 out of 33 layers and finetune the last 1 layer for 1 epoch. We use 1xT4 GPU vs. 1xP100 GPU for 1.7x speedup!\n* **Version 5** stay tuned for more versions...\n\n[1]: https://www.kaggle.com/code/yasufuminakama/fb3-deberta-v3-base-baseline-train\n[2]: https://www.kaggle.com/competitions/novozymes-enzyme-stability-prediction/discussion/356182\n[3]: https://www.kaggle.com/code/cdeotte/xgboost-5000-mutations-200-pdb-files-lb-0-40\n[4]: https://www.kaggle.com/competitions/novozymes-enzyme-stability-prediction/discussion/358320\n[5]: https://github.com/facebookresearch/esm\n[6]: https://www.kaggle.com/code/dschettler8845/novo-esp-fireprotdb-a-better-train-dataset","metadata":{"papermill":{"duration":0.011741,"end_time":"2022-10-29T04:55:04.612284","exception":false,"start_time":"2022-10-29T04:55:04.600543","status":"completed"},"tags":[]}},{"cell_type":"code","source":"! nvidia-smi","metadata":{"papermill":{"duration":1.167097,"end_time":"2022-10-29T04:55:05.78984","exception":false,"start_time":"2022-10-29T04:55:04.622743","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:08:56.362772Z","iopub.execute_input":"2023-06-23T20:08:56.363202Z","iopub.status.idle":"2023-06-23T20:08:58.204173Z","shell.execute_reply.started":"2023-06-23T20:08:56.363168Z","shell.execute_reply":"2023-06-23T20:08:58.201861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Directory settings","metadata":{"papermill":{"duration":0.011444,"end_time":"2022-10-29T04:55:05.81283","exception":false,"start_time":"2022-10-29T04:55:05.801386","status":"completed"},"tags":[]}},{"cell_type":"code","source":"VER = 1","metadata":{"papermill":{"duration":0.018537,"end_time":"2022-10-29T04:55:05.841152","exception":false,"start_time":"2022-10-29T04:55:05.822615","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:08:58.210258Z","iopub.execute_input":"2023-06-23T20:08:58.211466Z","iopub.status.idle":"2023-06-23T20:08:58.216707Z","shell.execute_reply.started":"2023-06-23T20:08:58.211424Z","shell.execute_reply":"2023-06-23T20:08:58.215467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nOUTPUT_DIR = f'VER_{VER}/'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)","metadata":{"papermill":{"duration":0.018289,"end_time":"2022-10-29T04:55:05.869066","exception":false,"start_time":"2022-10-29T04:55:05.850777","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:08:58.218571Z","iopub.execute_input":"2023-06-23T20:08:58.220311Z","iopub.status.idle":"2023-06-23T20:08:58.23122Z","shell.execute_reply.started":"2023-06-23T20:08:58.220274Z","shell.execute_reply":"2023-06-23T20:08:58.230085Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CFG","metadata":{"papermill":{"duration":0.009896,"end_time":"2022-10-29T04:55:05.889069","exception":false,"start_time":"2022-10-29T04:55:05.879173","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\nclass CFG:\n    wandb=True\n    competition='CAFA5'\n    _wandb_kernel='simon'\n    debug=False\n    apex=True\n    print_freq=20\n    num_workers=2  \n    #model=\"Rostlab/prot_bert\"\n    model=\"facebook/esm2_t6_8M_UR50D\" \n    gradient_checkpointing=False\n    scheduler='constant' # ['linear', 'cosine', 'constant']\n    batch_scheduler=True\n    num_warmup_steps=0\n    \n    # LEARNING RATE. \n    # Suggested: Prot_bert = 5e-6, ESM2 = 5e-5\n    epochs=1\n    num_cycles=1.0\n    encoder_lr=5e-5\n    decoder_lr=5e-5\n    batch_size=16\n    \n    # MODEL INFO - PROT_BERT or ONTO PROTEIN\n    total_layers = 30 \n    initial_layers = 5 \n    layers_per_block = 16 \n    # MODEL INFO - FACEBOOK ESM2\n    if 'esm2' in model:\n        total_layers = int(model.split('_')[1][1:])\n        initial_layers = 2 \n        layers_per_block = 16 # see @cdeotte comments \n        \n    # FREEZE\n    # Suggested: Prot_bert -8, ESM2 -1\n    num_freeze_layers = total_layers-8 # see @cdeotte comments\n    # NO FREEZE\n    #num_freeze_layers = 0\n    \n    min_lr=1e-6\n    eps=1e-6\n    betas=(0.9, 0.999)\n    max_len=512\n    weight_decay=0.01\n    gradient_accumulation_steps=1\n    max_grad_norm=1000\n    target_cols=None # SV: WILL SET THIS LATER\n    num_labels = 500\n    seed=42\n    n_fold=5\n    trn_fold=[0]#[0,1,2,3,4]\n    train=True\n    pca_dim = 64\n    \nif CFG.debug:\n    CFG.epochs = 1\n    CFG.trn_fold = [0]","metadata":{"papermill":{"duration":0.024977,"end_time":"2022-10-29T04:55:05.924044","exception":false,"start_time":"2022-10-29T04:55:05.899067","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:08:58.236463Z","iopub.execute_input":"2023-06-23T20:08:58.237903Z","iopub.status.idle":"2023-06-23T20:08:58.257382Z","shell.execute_reply.started":"2023-06-23T20:08:58.237871Z","shell.execute_reply":"2023-06-23T20:08:58.256238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# wandb\n# ====================================================\nif CFG.wandb:\n    \n    import wandb\n\n    try:\n        from kaggle_secrets import UserSecretsClient\n        user_secrets = UserSecretsClient()\n        secret_value_0 = user_secrets.get_secret(\"wandb-key\")\n        wandb.login(key=secret_value_0)\n        anony = None\n    except:\n        anony = \"must\"\n        print('If you want to use your W&B account, go to Add-ons -> Secrets and provide your W&B access token. Use the Label name as wandb_api. \\nGet your W&B access token from here: https://wandb.ai/authorize')\n\n\n    def class2dict(f):\n        return dict((name, getattr(f, name)) for name in dir(f) if not name.startswith('__'))\n\n    run = wandb.init(project='CAFA5', \n                     name=CFG.model,\n                     config=class2dict(CFG),\n                     group=CFG.model,\n                     job_type=\"train\",\n                     anonymous=anony)","metadata":{"papermill":{"duration":0.021665,"end_time":"2022-10-29T04:55:05.956519","exception":false,"start_time":"2022-10-29T04:55:05.934854","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:08:58.262723Z","iopub.execute_input":"2023-06-23T20:08:58.265156Z","iopub.status.idle":"2023-06-23T20:08:58.283452Z","shell.execute_reply.started":"2023-06-23T20:08:58.265124Z","shell.execute_reply":"2023-06-23T20:08:58.281971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Library","metadata":{"papermill":{"duration":0.010495,"end_time":"2022-10-29T04:55:05.977408","exception":false,"start_time":"2022-10-29T04:55:05.966913","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Library\n# ====================================================\nimport os\nimport gc\nimport re\nimport ast\nimport sys\nimport copy\nimport json\nimport time\nimport math\nimport string\nimport pickle\nimport random\nimport joblib\nimport itertools\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\npd.set_option('display.max_rows', 500)\npd.set_option('display.max_columns', 500)\npd.set_option('display.width', 1000)\nfrom tqdm.auto import tqdm\nfrom sklearn.metrics import mean_squared_error\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold\n\n#os.system('pip install iterative-stratification==0.1.7')\n#from iterstrat.ml_stratifiers import MultilabelStratifiedKFold\n\nimport torch\nimport torch.nn as nn\nfrom torch.nn import Parameter\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD, AdamW\nfrom torch.utils.data import DataLoader, Dataset\n\n#os.system('pip uninstall -y transformers')\n#os.system('pip uninstall -y tokenizers')\n#os.system('python -m pip install --no-index --find-links=../input/fb3-pip-wheels transformers')\n#os.system('python -m pip install --no-index --find-links=../input/fb3-pip-wheels tokenizers')\nos.system('pip install transformers --upgrade')\n#os.system('pip install tokenizers --upgrade')\n\nimport tokenizers\nimport transformers\nprint(f\"tokenizers.__version__: {tokenizers.__version__}\")\nprint(f\"transformers.__version__: {transformers.__version__}\")\nfrom transformers import AutoTokenizer, AutoModel, AutoConfig\nfrom transformers import get_linear_schedule_with_warmup, get_cosine_schedule_with_warmup\nfrom transformers import get_constant_schedule_with_warmup\n%env TOKENIZERS_PARALLELISM=true\n\n# CPMP: declare the two GPUs\nos.environ['CUDA_VISIBLE_DEVICES'] = \"0,1\"\n\n# CPMP: avoids some issues when using more than one worker\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"_kg_hide-output":true,"papermill":{"duration":27.103784,"end_time":"2022-10-29T04:55:33.092102","exception":false,"start_time":"2022-10-29T04:55:05.988318","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:08:58.291054Z","iopub.execute_input":"2023-06-23T20:08:58.293423Z","iopub.status.idle":"2023-06-23T20:09:33.769039Z","shell.execute_reply.started":"2023-06-23T20:08:58.293391Z","shell.execute_reply":"2023-06-23T20:09:33.767991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{"papermill":{"duration":0.015922,"end_time":"2022-10-29T04:55:33.126937","exception":false,"start_time":"2022-10-29T04:55:33.111015","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Utils\n# ====================================================\n\n# SV: WE ADJUSTED HERE TO USE ROC AUC\nfrom sklearn.metrics import roc_auc_score\n\ndef get_score(y_trues, y_preds):\n    # Convert logits to probabilities\n    y_preds = 1 / (1 + np.exp(-y_preds))\n    roc_auc_scores = []\n    for i in range(y_trues.shape[1]):\n        if len(np.unique(y_trues[:, i])) == 1:  # Skip if only one label is present\n            continue\n        roc_auc_scores.append(roc_auc_score(y_trues[:, i], y_preds[:, i]))\n    mcauc_score = np.mean(roc_auc_scores) if roc_auc_scores else None\n    return mcauc_score, roc_auc_scores\n\ndef get_logger(filename=OUTPUT_DIR+'train'):\n    from logging import getLogger, INFO, StreamHandler, FileHandler, Formatter\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=f\"{filename}.log\")\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = get_logger()\n\n\ndef seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    \nseed_everything(seed=42)","metadata":{"papermill":{"duration":0.036813,"end_time":"2022-10-29T04:55:33.179642","exception":false,"start_time":"2022-10-29T04:55:33.142829","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:16:23.202579Z","iopub.execute_input":"2023-06-23T20:16:23.202978Z","iopub.status.idle":"2023-06-23T20:16:23.216853Z","shell.execute_reply.started":"2023-06-23T20:16:23.202942Z","shell.execute_reply":"2023-06-23T20:16:23.215828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{"papermill":{"duration":0.016251,"end_time":"2022-10-29T04:55:33.212077","exception":false,"start_time":"2022-10-29T04:55:33.195826","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# SV: THIS IS NEW CODE. WE CREATE SIMILAR FORMAT AS IN THE OTHER COMPETITIONS\n\n# Prepare dataset\nfrom Bio import SeqIO\n\n\nprint(\"GENERATE ID AND SEQ LIST FOR TRAIN.\")\ntrain_ids = []\ntrain_sequences = []\nfor record in tqdm(SeqIO.parse(\"/kaggle/input/cafa-5-protein-function-prediction/Train/train_sequences.fasta\", \"fasta\")):\n    train_ids.append(record.id)\n    train_sequences.append(str(record.seq))\n    \n# Create the labels\nprint(\"GENERATE TARGETS FOR ENTRY IDS (\"+str(500)+\" MOST COMMON GO TERMS)\")\nids = train_ids\nlabels = pd.read_csv(\"/kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv\", sep = \"\\t\")\n\ntop_terms = labels.groupby(\"term\")[\"EntryID\"].count().sort_values(ascending=False)\nlabels_names = top_terms[:CFG.num_labels].index.values\ntrain_labels_sub = labels[(labels.term.isin(labels_names)) & (labels.EntryID.isin(ids))]\nid_labels = train_labels_sub.groupby('EntryID')['term'].apply(list).to_dict()\n\ngo_terms_map = {label: i for i, label in enumerate(labels_names)}\nlabels_matrix = np.empty((len(ids), len(labels_names)))\n\nfor index, id in tqdm(enumerate(ids)):\n    id_gos_list = id_labels[id]\n    temp = [go_terms_map[go] for go in labels_names if go in id_gos_list]\n    labels_matrix[index, temp] = 1\n    \n# now we can set the targets\nCFG.target_cols = labels_names\n\n# put the info in a dataframe with columns id, sequence, label_1, label_2, ...\ntrain = pd.DataFrame({'id': ids, 'sequence': train_sequences})\ntrain = pd.concat([train, pd.DataFrame(labels_matrix, columns=labels_names)], axis=1)\nprint(\"Train has shape: \", train.shape)\ntrain.head()","metadata":{"papermill":{"duration":0.163781,"end_time":"2022-10-29T04:55:33.392142","exception":false,"start_time":"2022-10-29T04:55:33.228361","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:09:33.791003Z","iopub.execute_input":"2023-06-23T20:09:33.791561Z","iopub.status.idle":"2023-06-23T20:10:56.105281Z","shell.execute_reply.started":"2023-06-23T20:09:33.791528Z","shell.execute_reply":"2023-06-23T20:10:56.104227Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def add_spaces(x):\n    return \" \".join(list(x))\ntrain.sequence = train.sequence.map(add_spaces)\nprint('Train has shape',train.shape)\ntrain.head()","metadata":{"papermill":{"duration":0.124195,"end_time":"2022-10-29T04:55:33.536943","exception":false,"start_time":"2022-10-29T04:55:33.412748","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:10:56.106394Z","iopub.execute_input":"2023-06-23T20:10:56.106754Z","iopub.status.idle":"2023-06-23T20:11:00.444257Z","shell.execute_reply.started":"2023-06-23T20:10:56.106719Z","shell.execute_reply":"2023-06-23T20:11:00.443251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV split","metadata":{"papermill":{"duration":0.018886,"end_time":"2022-10-29T04:55:34.022202","exception":false,"start_time":"2022-10-29T04:55:34.003316","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from sklearn.model_selection import GroupKFold","metadata":{"papermill":{"duration":0.030788,"end_time":"2022-10-29T04:55:34.072127","exception":false,"start_time":"2022-10-29T04:55:34.041339","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:11:00.449375Z","iopub.execute_input":"2023-06-23T20:11:00.450118Z","iopub.status.idle":"2023-06-23T20:11:00.455204Z","shell.execute_reply.started":"2023-06-23T20:11:00.45009Z","shell.execute_reply":"2023-06-23T20:11:00.454098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# CV split\n# ====================================================\nFold = GroupKFold(n_splits=CFG.n_fold) \nfor n, (train_index, val_index) in enumerate(Fold.split(train, train[CFG.target_cols], train.id)):\n    train.loc[val_index, 'fold'] = int(n)\ntrain['fold'] = train['fold'].astype(int)\ndisplay(train.groupby('fold').size())","metadata":{"papermill":{"duration":0.050009,"end_time":"2022-10-29T04:55:34.142355","exception":false,"start_time":"2022-10-29T04:55:34.092346","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:11:00.45689Z","iopub.execute_input":"2023-06-23T20:11:00.45757Z","iopub.status.idle":"2023-06-23T20:11:01.654617Z","shell.execute_reply.started":"2023-06-23T20:11:00.457535Z","shell.execute_reply":"2023-06-23T20:11:01.653693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.debug:\n    display(train.groupby('fold').size())\n    train = train.sample(n=1000, random_state=0).reset_index(drop=True)\n    display(train.groupby('fold').size())","metadata":{"papermill":{"duration":0.033455,"end_time":"2022-10-29T04:55:34.197662","exception":false,"start_time":"2022-10-29T04:55:34.164207","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:11:01.65623Z","iopub.execute_input":"2023-06-23T20:11:01.656663Z","iopub.status.idle":"2023-06-23T20:11:01.686449Z","shell.execute_reply.started":"2023-06-23T20:11:01.656626Z","shell.execute_reply":"2023-06-23T20:11:01.685474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Tokenizer","metadata":{"papermill":{"duration":0.019411,"end_time":"2022-10-29T04:55:34.238008","exception":false,"start_time":"2022-10-29T04:55:34.218597","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# tokenizer\n# ====================================================\ntokenizer = AutoTokenizer.from_pretrained(CFG.model)\ntokenizer.save_pretrained(OUTPUT_DIR+'tokenizer/')\nCFG.tokenizer = tokenizer","metadata":{"papermill":{"duration":1.328653,"end_time":"2022-10-29T04:55:35.585806","exception":false,"start_time":"2022-10-29T04:55:34.257153","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:11:01.687998Z","iopub.execute_input":"2023-06-23T20:11:01.688353Z","iopub.status.idle":"2023-06-23T20:11:02.150849Z","shell.execute_reply.started":"2023-06-23T20:11:01.688321Z","shell.execute_reply":"2023-06-23T20:11:02.149828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tokenizer","metadata":{"papermill":{"duration":0.021444,"end_time":"2022-10-29T04:55:35.619225","exception":false,"start_time":"2022-10-29T04:55:35.597781","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:11:02.152613Z","iopub.execute_input":"2023-06-23T20:11:02.153325Z","iopub.status.idle":"2023-06-23T20:11:02.160865Z","shell.execute_reply.started":"2023-06-23T20:11:02.153286Z","shell.execute_reply":"2023-06-23T20:11:02.159821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{"papermill":{"duration":0.011968,"end_time":"2022-10-29T04:55:35.642685","exception":false,"start_time":"2022-10-29T04:55:35.630717","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Dataset\n# ====================================================\ndef prepare_input(cfg, text):\n    inputs = cfg.tokenizer.encode_plus(\n        text, \n        return_tensors=None, \n        add_special_tokens=True, \n        max_length=cfg.max_len,\n        pad_to_max_length=True,\n        truncation=True\n    )\n    for k, v in inputs.items():\n        inputs[k] = torch.tensor(v, dtype=torch.long)\n    return inputs\n\n\nclass TrainDataset(Dataset):\n    def __init__(self, cfg, df):\n        self.cfg = cfg\n        # SV: ONLY HAVE ONE SEQUENCE PER SAMPLE\n        self.texts1 = df['sequence'].values\n        self.labels = df[cfg.target_cols].values\n\n    def __len__(self):\n        return len(self.texts1)\n\n    def __getitem__(self, item):\n        inputs1 = prepare_input(self.cfg, self.texts1[item])\n        labels = torch.tensor(self.labels[item], dtype=torch.float)\n        \n        # SV: ADJUSTED HERE AS WE ONLY HAVE ONE INPUT.\n        # CPMP: return a single dictionary containing all inputs to avoid issues when using DP\n        return {'input1_ids' : inputs1['input_ids'], \n                'input1_attention_mask' : inputs1['attention_mask'], \n                'labels' : labels}\n    \n\ndef collate(inputs):\n    mask_len = int(inputs[\"attention_mask\"].sum(axis=1).max())\n    for k, v in inputs.items():\n        # CPMP: no need to truncate labels\n        if k != 'labels':\n            inputs[k] = inputs[k][:,:mask_len]\n    return inputs","metadata":{"papermill":{"duration":0.028004,"end_time":"2022-10-29T04:55:35.682761","exception":false,"start_time":"2022-10-29T04:55:35.654757","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:11:02.162667Z","iopub.execute_input":"2023-06-23T20:11:02.16351Z","iopub.status.idle":"2023-06-23T20:11:02.17807Z","shell.execute_reply.started":"2023-06-23T20:11:02.163451Z","shell.execute_reply":"2023-06-23T20:11:02.176966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EDA Dataloader\nBelow displays the output of our data loader using a fake data example. Our data loader provides us with \"two sentences\" (i.e. both the wild type sequence and mutation sequence) and it provides us with the location of the mutation. Our model will use all this information. Note in the output that the first token id is `<cls>` and the last token id is `<eos>` and the tokens inbetween are the amino acids.","metadata":{"papermill":{"duration":0.012106,"end_time":"2022-10-29T04:55:35.707157","exception":false,"start_time":"2022-10-29T04:55:35.695051","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# SV: ADJUSTED DUMMY EXAMPLE\nfake = pd.DataFrame(columns=['sequence'])\nfake['sequence'] = ['V P V N P E','V P V N P E']\nfake[CFG.target_cols] = [1 if x < 250 else 0 for x in range(500)]\nfake.head()","metadata":{"papermill":{"duration":0.031033,"end_time":"2022-10-29T04:55:35.750689","exception":false,"start_time":"2022-10-29T04:55:35.719656","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:11:02.180023Z","iopub.execute_input":"2023-06-23T20:11:02.180736Z","iopub.status.idle":"2023-06-23T20:11:02.549908Z","shell.execute_reply.started":"2023-06-23T20:11:02.180702Z","shell.execute_reply":"2023-06-23T20:11:02.548474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# SV: ADJUSTED DUMMY EXAMPLE\nclass CFG_fake:\n    max_len = 8\n    target_cols = CFG.target_cols\n    tokenizer = CFG.tokenizer\n    \ntrain_dataset = TrainDataset(CFG_fake, fake)\ntrain_loader = DataLoader(train_dataset,\n                          batch_size=2,\n                          shuffle=False)\nfor batch in train_loader:\n    break\n    \nbatch","metadata":{"papermill":{"duration":0.034639,"end_time":"2022-10-29T04:55:35.798299","exception":false,"start_time":"2022-10-29T04:55:35.76366","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:11:02.551528Z","iopub.execute_input":"2023-06-23T20:11:02.552143Z","iopub.status.idle":"2023-06-23T20:11:02.67444Z","shell.execute_reply.started":"2023-06-23T20:11:02.552106Z","shell.execute_reply":"2023-06-23T20:11:02.672838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model\nThe model architecture takes both wild type sequence and mutation sequence and location of mutation as input. Then it subtracts their embeddings and concatenates that with their individual embeddings. It does this with both single mutation token position and mean pooling of entire sequence. Finally a dense layer makes the regression prediction. See the model architecture below.","metadata":{"papermill":{"duration":0.012016,"end_time":"2022-10-29T04:55:35.822305","exception":false,"start_time":"2022-10-29T04:55:35.810289","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Model\n# ====================================================\nclass MeanPooling(nn.Module):\n    def __init__(self):\n        super(MeanPooling, self).__init__()\n        \n    def forward(self, last_hidden_state, attention_mask):\n        input_mask_expanded = attention_mask.unsqueeze(-1).expand(last_hidden_state.size()).float()\n        sum_embeddings = torch.sum(last_hidden_state * input_mask_expanded, 1)\n        sum_mask = input_mask_expanded.sum(1)\n        sum_mask = torch.clamp(sum_mask, min=1e-9)\n        mean_embeddings = sum_embeddings / sum_mask\n        return mean_embeddings\n    \n\nclass CustomModel(nn.Module):\n    def __init__(self, cfg, config_path=None, pretrained=False):\n        super().__init__()\n        self.cfg = cfg\n        if config_path is None:\n            self.config = AutoConfig.from_pretrained(cfg.model, output_hidden_states=True)\n            self.config.hidden_dropout = 0.\n            self.config.hidden_dropout_prob = 0.\n            self.config.attention_dropout = 0.\n            self.config.attention_probs_dropout_prob = 0.\n            LOGGER.info(self.config)\n        else:\n            self.config = torch.load(config_path)\n            #self.config = AutoConfig.from_pretrained(config_path)\n        if pretrained:\n            self.model = AutoModel.from_pretrained(cfg.model, config=self.config)\n        else:\n            self.model = AutoModel.from_config(self.config)\n            \n        if self.cfg.gradient_checkpointing:\n            self.model.gradient_checkpointing_enable()\n        self.pool = MeanPooling()\n        self.fc1 = nn.Linear(self.config.hidden_size, CFG.num_labels)\n        # SV: ADJUST\n        self._init_weights(self.fc1)\n        \n    def _init_weights(self, module):\n        if isinstance(module, nn.Linear):\n            module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)\n            if module.bias is not None:\n                module.bias.data.zero_()\n        elif isinstance(module, nn.Embedding):\n            module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)\n            if module.padding_idx is not None:\n                module.weight.data[module.padding_idx].zero_()\n        elif isinstance(module, nn.LayerNorm):\n            module.bias.data.zero_()\n            module.weight.data.fill_(1.0)\n        \n    def feature(self, input_ids, attention_mask):\n        # SV: ADJUST\n        # CPMP: pass inputs explicitly\n        outputs = self.model(input_ids=input_ids, attention_mask=attention_mask)\n        last_hidden_states = outputs[0]\n        feature = self.pool(last_hidden_states, attention_mask)\n        return feature\n\n    def forward(self, batch):\n        # SV: ADJUSTED FOR OUR EXAMPLE\n        # CPMP: change code to read inputs from a single dictionary\n        output = self.fc1(self.feature(batch['input1_ids'], batch['input1_attention_mask']))\n        return output","metadata":{"papermill":{"duration":0.032387,"end_time":"2022-10-29T04:55:35.866761","exception":false,"start_time":"2022-10-29T04:55:35.834374","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:11:02.677316Z","iopub.execute_input":"2023-06-23T20:11:02.677599Z","iopub.status.idle":"2023-06-23T20:11:02.695541Z","shell.execute_reply.started":"2023-06-23T20:11:02.677568Z","shell.execute_reply":"2023-06-23T20:11:02.694615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss","metadata":{"papermill":{"duration":0.011284,"end_time":"2022-10-29T04:55:35.889794","exception":false,"start_time":"2022-10-29T04:55:35.87851","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Loss\n# ====================================================\n# SV: ADJUST\nclass MultiLabelClassificationLoss(nn.Module):\n    def __init__(self, reduction='mean'):\n        super().__init__()\n        self.bce_with_logits = nn.BCEWithLogitsLoss(reduction=reduction)\n\n    def forward(self, y_pred, y_true):\n        return self.bce_with_logits(y_pred, y_true)","metadata":{"papermill":{"duration":0.022109,"end_time":"2022-10-29T04:55:35.923845","exception":false,"start_time":"2022-10-29T04:55:35.901736","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:11:02.697498Z","iopub.execute_input":"2023-06-23T20:11:02.698679Z","iopub.status.idle":"2023-06-23T20:11:02.710198Z","shell.execute_reply.started":"2023-06-23T20:11:02.698645Z","shell.execute_reply":"2023-06-23T20:11:02.70918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helpler functions","metadata":{"papermill":{"duration":0.012573,"end_time":"2022-10-29T04:55:35.948018","exception":false,"start_time":"2022-10-29T04:55:35.935445","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Helper functions\n# ====================================================\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\ndef asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\n\ndef timeSince(since, percent):\n    now = time.time()\n    s = now - since\n    es = s / (percent)\n    rs = es - s\n    return '%s (remain %s)' % (asMinutes(s), asMinutes(rs))\n\n\ndef train_fn(fold, train_loader, model, criterion, optimizer, epoch, scheduler, device):\n    model.train()\n    scaler = torch.cuda.amp.GradScaler(enabled=CFG.apex)\n    losses = AverageMeter()\n    start = end = time.time()\n    global_step = 0\n    \n    # CPMP: iterating on batch dictionaries\n    for step, batch in enumerate(train_loader):\n        for k, v in batch.items():\n            batch[k] = v.to(device)\n        batch_size = batch['labels'].size(0)\n        with torch.cuda.amp.autocast(enabled=CFG.apex):\n            y_preds = model(batch)\n            loss = criterion(y_preds, batch['labels'])\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        losses.update(loss.item(), batch_size)\n        scaler.scale(loss).backward()\n        grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.max_grad_norm)\n        if (step + 1) % CFG.gradient_accumulation_steps == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            global_step += 1\n            if CFG.batch_scheduler:\n                scheduler.step()\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(train_loader)-1):\n            print('Epoch: [{0}][{1}/{2}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  'Grad: {grad_norm:.4f}  '\n                  'LR: {lr:.8f}  '\n                  .format(epoch+1, step, len(train_loader), \n                          remain=timeSince(start, float(step+1)/len(train_loader)),\n                          loss=losses,\n                          grad_norm=grad_norm,\n                          lr=scheduler.get_lr()[0]))\n        if CFG.wandb:\n            wandb.log({f\"[fold{fold}] loss\": losses.val,\n                       f\"[fold{fold}] lr\": scheduler.get_lr()[0]})\n    return losses.avg\n\n\ndef valid_fn(valid_loader, model, criterion, device):\n    losses = AverageMeter()\n    model.eval()\n    preds = []\n    start = end = time.time()\n    # CPMP: iterating on batch dictionaries\n    for step, batch in enumerate(valid_loader):\n        for k, v in batch.items():\n            batch[k] = v.to(device)\n        batch_size = batch['labels'].size(0)\n        with torch.no_grad():\n            y_preds = model(batch)\n            loss = criterion(y_preds, batch['labels'])\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        losses.update(loss.item(), batch_size)\n        preds.append(y_preds.to('cpu').numpy())\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(valid_loader)-1):\n            print('EVAL: [{0}/{1}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  .format(step, len(valid_loader),\n                          loss=losses,\n                          remain=timeSince(start, float(step+1)/len(valid_loader))))\n    predictions = np.concatenate(preds)\n    return losses.avg, predictions","metadata":{"papermill":{"duration":0.035075,"end_time":"2022-10-29T04:55:35.995055","exception":false,"start_time":"2022-10-29T04:55:35.95998","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:11:02.711463Z","iopub.execute_input":"2023-06-23T20:11:02.711965Z","iopub.status.idle":"2023-06-23T20:11:02.733178Z","shell.execute_reply.started":"2023-06-23T20:11:02.711933Z","shell.execute_reply":"2023-06-23T20:11:02.732141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train loop","metadata":{"papermill":{"duration":0.011903,"end_time":"2022-10-29T04:55:36.018734","exception":false,"start_time":"2022-10-29T04:55:36.006831","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# train loop\n# ====================================================\ndef train_loop(folds, fold):\n    \n    LOGGER.info(f\"========== fold: {fold} training ==========\")\n\n    # ====================================================\n    # loader\n    # ====================================================\n    train_folds = folds[folds['fold'] != fold].reset_index(drop=True)\n    valid_folds = folds[folds['fold'] == fold].reset_index(drop=True)\n    valid_labels = valid_folds[CFG.target_cols].values\n    print('### train shape:',train_folds.shape)\n    \n    train_dataset = TrainDataset(CFG, train_folds)\n    valid_dataset = TrainDataset(CFG, valid_folds)\n\n    train_loader = DataLoader(train_dataset,\n                              batch_size=CFG.batch_size,\n                              shuffle=True,\n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=True)\n    valid_loader = DataLoader(valid_dataset,\n                              batch_size=CFG.batch_size * 2,\n                              shuffle=False,\n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n\n    # ====================================================\n    # model & optimizer\n    # ====================================================\n    model = CustomModel(CFG, config_path=None, pretrained=True)\n    \n    torch.save(model.config, OUTPUT_DIR+'config.pth')\n    \n    # FREEZE LAYERS\n    if CFG.num_freeze_layers>0:\n        print(f'### Freezing first {CFG.num_freeze_layers} layers.',\n              f'Leaving {CFG.total_layers-CFG.num_freeze_layers} layers unfrozen')\n        for name, param in list(model.named_parameters())\\\n            [:CFG.initial_layers+CFG.layers_per_block*CFG.num_freeze_layers]:     \n                param.requires_grad = False\n    model.to(device)\n    # CPMP: wrap the model to use all GPUs\n    model = nn.DataParallel(model)\n    \n    def get_optimizer_params(model, encoder_lr, decoder_lr, weight_decay=0.0):\n        param_optimizer = list(model.named_parameters())\n        no_decay = [\"bias\", \"LayerNorm.bias\", \"LayerNorm.weight\"]\n        # CPMP: when using DataParallel the original model is now the module attribute\n        optimizer_parameters = [\n            {'params': [p for n, p in model.module.model.named_parameters() if not any(nd in n for nd in no_decay)],\n             'lr': encoder_lr, 'weight_decay': weight_decay},\n            {'params': [p for n, p in model.module.model.named_parameters() if any(nd in n for nd in no_decay)],\n             'lr': encoder_lr, 'weight_decay': 0.0},\n            {'params': [p for n, p in model.module.named_parameters() if \"model\" not in n],\n             'lr': decoder_lr, 'weight_decay': 0.0}\n        ]\n        return optimizer_parameters\n\n    optimizer_parameters = get_optimizer_params(model,\n                                                encoder_lr=CFG.encoder_lr, \n                                                decoder_lr=CFG.decoder_lr,\n                                                weight_decay=CFG.weight_decay)\n    optimizer = AdamW(optimizer_parameters, lr=CFG.encoder_lr, eps=CFG.eps, betas=CFG.betas)\n    \n    # ====================================================\n    # scheduler\n    # ====================================================\n    def get_scheduler(cfg, optimizer, num_train_steps):\n        if cfg.scheduler == 'linear':\n            scheduler = get_linear_schedule_with_warmup(\n                optimizer, num_warmup_steps=cfg.num_warmup_steps, num_training_steps=num_train_steps\n            )\n        elif cfg.scheduler == 'constant':\n            scheduler = get_constant_schedule_with_warmup(\n                optimizer, num_warmup_steps=cfg.num_warmup_steps\n            )\n        elif cfg.scheduler == 'cosine':\n            scheduler = get_cosine_schedule_with_warmup(\n                optimizer, num_warmup_steps=cfg.num_warmup_steps, num_training_steps=num_train_steps, num_cycles=cfg.num_cycles\n            )\n        return scheduler\n    \n    num_train_steps = int(len(train_folds) / CFG.batch_size * CFG.epochs)\n    scheduler = get_scheduler(CFG, optimizer, num_train_steps)\n\n    # ====================================================\n    # loop\n    # ====================================================\n    criterion = MultiLabelClassificationLoss(reduction=\"mean\") #nn.SmoothL1Loss(reduction='mean')\n    \n    best_score = np.inf\n\n    for epoch in range(CFG.epochs):\n\n        start_time = time.time()\n\n        # train\n        avg_loss = train_fn(fold, train_loader, model, criterion, optimizer, epoch, scheduler, device)\n\n        # eval\n        avg_val_loss, predictions = valid_fn(valid_loader, model, criterion, device)\n        \n        # scoring\n        score, scores = get_score(valid_labels, predictions)\n\n        elapsed = time.time() - start_time\n\n        LOGGER.info(f'Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  avg_val_loss: {avg_val_loss:.4f}  time: {elapsed:.0f}s')\n        LOGGER.info(f'Epoch {epoch+1} - Score: {score:.4f}  Scores: {scores}')\n        if CFG.wandb:\n            wandb.log({f\"[fold{fold}] epoch\": epoch+1, \n                       f\"[fold{fold}] avg_train_loss\": avg_loss, \n                       f\"[fold{fold}] avg_val_loss\": avg_val_loss,\n                       f\"[fold{fold}] score\": score})\n        \n        if best_score > score:\n            best_score = score\n            LOGGER.info(f'Epoch {epoch+1} - Save Best Score: {best_score:.4f} Model')\n            # CPMP: save the original model. It is stored as the module attribute of the DP model.\n            torch.save({'model': model.module.state_dict(),\n                        'predictions': predictions},\n                        OUTPUT_DIR+f\"{CFG.model.replace('/', '-')}_fold{fold}_best.pth\")\n\n    predictions = torch.load(OUTPUT_DIR+f\"{CFG.model.replace('/', '-')}_fold{fold}_best.pth\", \n                             map_location=torch.device('cpu'))['predictions']\n    valid_folds[[f\"pred_{c}\" for c in CFG.target_cols]] = predictions\n\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    return valid_folds","metadata":{"papermill":{"duration":0.036756,"end_time":"2022-10-29T04:55:36.067149","exception":false,"start_time":"2022-10-29T04:55:36.030393","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:11:02.735003Z","iopub.execute_input":"2023-06-23T20:11:02.735756Z","iopub.status.idle":"2023-06-23T20:11:02.760179Z","shell.execute_reply.started":"2023-06-23T20:11:02.735722Z","shell.execute_reply":"2023-06-23T20:11:02.759157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == '__main__':\n    \n    def get_result(oof_df):\n        labels = oof_df[CFG.target_cols].values\n        preds = oof_df[[f\"pred_{c}\" for c in CFG.target_cols]].values\n        score, scores = get_score(labels, preds)\n        LOGGER.info(f'Score: {score:<.4f}  Scores: {scores}')\n    \n    if CFG.train:\n        oof_df = pd.DataFrame()\n        for fold in range(CFG.n_fold):\n            if fold in CFG.trn_fold:\n                _oof_df = train_loop(train, fold)\n                oof_df = pd.concat([oof_df, _oof_df])\n                LOGGER.info(f\"========== fold: {fold} result ==========\")\n                get_result(_oof_df)\n        oof_df = oof_df.reset_index(drop=True)\n        LOGGER.info(f\"========== CV ==========\")\n        get_result(oof_df)\n        oof_df.to_pickle(OUTPUT_DIR+'oof_df.pkl')\n        \n    if CFG.wandb:\n        wandb.finish()","metadata":{"papermill":{"duration":5702.136941,"end_time":"2022-10-29T06:30:38.216293","exception":false,"start_time":"2022-10-29T04:55:36.079352","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:16:28.853838Z","iopub.execute_input":"2023-06-23T20:16:28.854203Z","iopub.status.idle":"2023-06-23T20:16:37.508119Z","shell.execute_reply.started":"2023-06-23T20:16:28.854174Z","shell.execute_reply":"2023-06-23T20:16:37.507053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Compute OOF Score","metadata":{"papermill":{"duration":0.020677,"end_time":"2022-10-29T06:30:38.257522","exception":false,"start_time":"2022-10-29T06:30:38.236845","status":"completed"},"tags":[]}},{"cell_type":"code","source":"oof_df.head()","metadata":{"papermill":{"duration":0.062218,"end_time":"2022-10-29T06:30:38.339865","exception":false,"start_time":"2022-10-29T06:30:38.277647","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:16:56.558361Z","iopub.execute_input":"2023-06-23T20:16:56.558726Z","iopub.status.idle":"2023-06-23T20:16:57.021274Z","shell.execute_reply.started":"2023-06-23T20:16:56.558696Z","shell.execute_reply":"2023-06-23T20:16:57.018394Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Infer Test Data","metadata":{"papermill":{"duration":0.019147,"end_time":"2022-10-29T06:30:38.599673","exception":false,"start_time":"2022-10-29T06:30:38.580526","status":"completed"},"tags":[]}},{"cell_type":"code","source":"CFG.path = OUTPUT_DIR\nCFG.config_path = CFG.path+'config.pth'","metadata":{"papermill":{"duration":0.027692,"end_time":"2022-10-29T06:30:38.646991","exception":false,"start_time":"2022-10-29T06:30:38.619299","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:16:58.839571Z","iopub.execute_input":"2023-06-23T20:16:58.840378Z","iopub.status.idle":"2023-06-23T20:16:58.844517Z","shell.execute_reply.started":"2023-06-23T20:16:58.840343Z","shell.execute_reply":"2023-06-23T20:16:58.843603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# SV: ADJUST\nclass TestDataset(Dataset):\n    def __init__(self, cfg, df):\n        self.cfg = cfg\n        self.texts1 = df['sequence'].values\n\n    def __len__(self):\n        return len(self.texts1)\n\n    def __getitem__(self, item):\n        inputs1 = prepare_input(self.cfg, self.texts1[item])\n        return inputs1\n    \n    def __getitem__(self, item):\n        inputs1 = prepare_input(self.cfg, self.texts1[item])\n        # CPMP: return a single diciotnary containing all inputs\n        return {'input1_ids' : inputs1['input_ids'], \n                'input1_attention_mask' : inputs1['attention_mask']}","metadata":{"papermill":{"duration":0.032845,"end_time":"2022-10-29T06:30:38.69946","exception":false,"start_time":"2022-10-29T06:30:38.666615","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:16:59.559655Z","iopub.execute_input":"2023-06-23T20:16:59.560318Z","iopub.status.idle":"2023-06-23T20:16:59.567432Z","shell.execute_reply.started":"2023-06-23T20:16:59.560282Z","shell.execute_reply":"2023-06-23T20:16:59.566212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# inference\n# ====================================================\ndef inference_fn(test_loader, model, device):\n    preds = []\n    model.eval()\n    model.to(device)\n    # CPMP: using 2 GPU\n    model = nn.DataParallel(model)\n    tk0 = tqdm(test_loader, total=len(test_loader))\n    # CPMP: iterating on batch dictionaries\n    for batch in tk0:\n        for k, v in batch.items():\n            batch[k] = v.to(device)\n        with torch.no_grad():\n            y_preds = model(batch)\n        # ADJUST TO PUTPUT PROBABILITIES\n        preds.append(torch.sigmoid(y_preds).to('cpu').numpy())\n    predictions = np.concatenate(preds)\n    return predictions","metadata":{"papermill":{"duration":0.030252,"end_time":"2022-10-29T06:30:38.749663","exception":false,"start_time":"2022-10-29T06:30:38.719411","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:17:00.422753Z","iopub.execute_input":"2023-06-23T20:17:00.423746Z","iopub.status.idle":"2023-06-23T20:17:00.431225Z","shell.execute_reply.started":"2023-06-23T20:17:00.423705Z","shell.execute_reply":"2023-06-23T20:17:00.430226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# SV: ADJUST\n# Prepare dataset\nfrom Bio import SeqIO\n\n\nprint(\"GENERATE ID AND SEQ LIST FOR TRAIN.\")\ntest_ids = []\ntest_sequences = []\nfor record in tqdm(SeqIO.parse(\"/kaggle/input/cafa-5-protein-function-prediction/Test (Targets)/testsuperset.fasta\", \"fasta\")):\n    test_ids.append(record.id)\n    test_sequences.append(str(record.seq))\n\n# put the info in a dataframe with columns id, sequence, label_1, label_2, ...\ntest = pd.DataFrame({'id': test_ids, 'sequence': test_sequences})\ntest.sequence = test.sequence.map(add_spaces)\ntest.head()","metadata":{"papermill":{"duration":0.031679,"end_time":"2022-10-29T06:30:38.801503","exception":false,"start_time":"2022-10-29T06:30:38.769824","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:17:01.650668Z","iopub.execute_input":"2023-06-23T20:17:01.65137Z","iopub.status.idle":"2023-06-23T20:17:04.115673Z","shell.execute_reply.started":"2023-06-23T20:17:01.651335Z","shell.execute_reply":"2023-06-23T20:17:04.11461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = TestDataset(CFG, test)\ntest_loader = DataLoader(test_dataset,\n                         batch_size=CFG.batch_size,\n                         shuffle=False,\n                         #collate_fn=DataCollatorWithPadding(tokenizer=CFG.tokenizer, padding='longest'),\n                         num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\npredictions_ = []\nfor fold in CFG.trn_fold:\n    model = CustomModel(CFG, config_path=CFG.config_path, pretrained=False)\n    state = torch.load(CFG.path+f\"{CFG.model.replace('/', '-')}_fold{fold}_best.pth\",\n                       map_location=torch.device('cpu'))\n    model.load_state_dict(state['model'])\n    prediction = inference_fn(test_loader, model, device)\n    predictions_.append(prediction)\n    del model, state, prediction; gc.collect()\n    torch.cuda.empty_cache()\npredictions = np.mean(predictions_, axis=0)","metadata":{"papermill":{"duration":5877.596491,"end_time":"2022-10-29T08:08:40.324203","exception":false,"start_time":"2022-10-29T06:30:42.727712","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:17:20.587742Z","iopub.execute_input":"2023-06-23T20:17:20.588829Z","iopub.status.idle":"2023-06-23T20:26:09.392583Z","shell.execute_reply.started":"2023-06-23T20:17:20.588752Z","shell.execute_reply":"2023-06-23T20:26:09.391404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Assuming 'predictions' is a numpy array containing your model's output,\n# and 'labels_name' is a list containing the names of your labels.\n\n# Create a DataFrame for the predictions\npreds_df = pd.DataFrame(predictions, columns=CFG.target_cols)\n\n# Ensure that the original DataFrame and the predictions DataFrame are aligned.\nassert len(test) == len(preds_df)\n\n# Join the original DataFrame with the predictions DataFrame\nresult = pd.concat([test, preds_df], axis=1).drop(\"sequence\", axis=1)","metadata":{"papermill":{"duration":0.097905,"end_time":"2022-10-29T08:08:40.444076","exception":false,"start_time":"2022-10-29T08:08:40.346171","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T20:28:46.040338Z","iopub.execute_input":"2023-06-23T20:28:46.042313Z","iopub.status.idle":"2023-06-23T20:28:46.497713Z","shell.execute_reply.started":"2023-06-23T20:28:46.042272Z","shell.execute_reply":"2023-06-23T20:28:46.496729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import polars as pl\n\ndf = pl.from_pandas(result)\n# Assuming result is your DataFrame\ndf = df.melt(id_vars='id', variable_name='LABEL', value_name='probability')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-23T20:28:55.992988Z","iopub.execute_input":"2023-06-23T20:28:55.994144Z","iopub.status.idle":"2023-06-23T20:28:59.863707Z","shell.execute_reply.started":"2023-06-23T20:28:55.994107Z","shell.execute_reply":"2023-06-23T20:28:59.859732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.write_csv('submission.tsv', separator='\\t', has_header=False)","metadata":{},"execution_count":null,"outputs":[]}]}