{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.8.16","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":41875,"databundleVersionId":5521661}],"dockerImageVersionId":30474,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"\nTODO\n- Verify are in_axis_resources correct?\n- Try adding QuickGO data. Use protein relation graph for better embeddings?\n- Try \"Correct and label\" graph\n- Different loss functions\n- Benchmark a baseline without backbone state\n- Train the full backbone on 512 seq len\n- Try (1, 8) mesh size if this works\n- Try different less memory optimizer\n\n\nBENCHMARKING\n- Benchmark the time difference with and without the classifier\n- Benchmark no gradient \n- Benckmark just training a few layers\n- JAX PROFILER\n\"\"\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Installations, Setup and Imports\n!pip install -q omegaconf\n!pip install -q transformers datasets\n!pip install -q biopython\n\n# Commonly Used Libraries\nimport pandas as pd\nimport numpy as np\nfrom pathlib import Path\nimport collections\nimport termcolor\nimport functools \nimport random\nimport os\nimport re\n\nfrom tqdm.auto import tqdm\ntqdm.pandas()\n\nimport omegaconf\nimport wandb\n!wandb login '3b335317f20548af7e3b941d09a6de9f1736bd8d'\n\nimport transformers\nimport datasets\nimport sklearn\nimport sklearn.metrics\n\n\nfrom IPython.core.magic import register_line_cell_magic\n@register_line_cell_magic\ndef hyperparameters(hp_var_name, cell):\n    with open('experiment.yaml', 'w') as f:\n        f.write(cell)\n    HP = omegaconf.OmegaConf.load('experiment.yaml')\n    get_ipython().user_ns[hp_var_name] = HP","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q git+https://github.com/irhum/esmjax\n\n# JAX Imports #\nimport jax\nimport flax \nimport optax\n\nfrom flax.linen import partitioning as nn_partitioning\nfrom jax.experimental import maps, PartitionSpec as P, pjit","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%hyperparameters args\n\nbackbone_name: 'facebook/esm2_t48_15B_UR50D'\nmax_seq_len: 1024 # ~90% coverage @ 1024\nbatch_size: 256 # 256\ndebug: False","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"wandb.restore('protein_df.csv', 'uncategorized/runs/ibtdzsot')\ntokenizer = transformers.AutoTokenizer.from_pretrained(args.backbone_name)\n# args.max_seq_len = 4 if args.debug else args.max_seq_len\n\nIMP_RESIDUES = ['M', 'N', 'S', 'V', 'T', 'H', 'A', 'P', 'Y', 'I', 'D', 'W', 'E', 'Q', 'L', 'F', 'R', 'K', 'C', 'G']","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"protein_df = pd.read_csv('protein_df.csv')\n\nterm_id_to_term_name = protein_df.set_index('term_id').to_dict()['term']\nONTOLOGY_TERM_NAMES = [term_id_to_term_name[term_id] for term_id in range(protein_df.term_id.nunique())]\nNUM_ONTOLOGY_LABELS = len(ONTOLOGY_TERM_NAMES)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n\nfrom Bio import SeqIO\nimport collections\n\nprotein_id_to_sequence = {}\n\ncomp_dir = Path('/kaggle/input/cafa-5-protein-function-prediction')\ntrain_sequences = SeqIO.parse(comp_dir/'Train'/'train_sequences.fasta', 'fasta')\ntest_sequences = SeqIO.parse(comp_dir/'Test (Targets)'/'testsuperset.fasta', 'fasta')\n\nfor train_seq in train_sequences:\n    protein_id_to_sequence[train_seq.id] = ' '.join(list(str(train_seq.seq)))\n\ntest_protein_ids = []\nfor test_seq in tqdm(test_sequences):\n    protein_id_to_sequence[test_seq.id] = ' '.join(list(str(test_seq.seq)))\n    test_protein_ids.append(test_seq.id)\n    if args.debug and len(test_protein_ids) > 10:\n        break\n\nprotein_id_to_term_ids = collections.defaultdict(list)\nfor protein_id, term_id in tqdm(zip(protein_df.EntryID.values, protein_df.term_id.values), total=len(protein_df)):\n    protein_id_to_term_ids[protein_id].append(int(term_id))\n    \n# Convert list of term ids to string due to huggingface bug\nprotein_id_to_term_ids = {protein_id: '[SEP]'.join([str(tid) for tid in term_ids]) for protein_id, term_ids in protein_id_to_term_ids.items()}","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\ndef process_protein_sequence(protein_id):\n    protein_sequence = protein_id_to_sequence[protein_id]\n    \n    tokenized_input = tokenizer(\n        protein_sequence,\n        add_special_tokens=True,\n        max_length=args.max_seq_len,\n        padding='max_length',\n        truncation=True,\n    )\n    protein_input_ids, attention_mask = tokenized_input['input_ids'], tokenized_input['attention_mask']\n    \n    output_dict = {\n        'input_ids': protein_input_ids,\n        'attention_mask': attention_mask,\n        'protein_id': protein_id,\n    }\n    \n    # For training and validation set, also add term_ids\n    if protein_id in protein_id_to_term_ids:\n        term_ids = protein_id_to_term_ids[protein_id].split('[SEP]')\n        term_ids = [int(term_id) for term_id in term_ids]\n        output_dict['term_ids'] = term_ids\n    else:\n        output_dict['term_ids'] = [-100]\n    \n    return output_dict\n\n\ntrain_protein_ids = list(protein_df.EntryID.unique())\ntrain_protein_ids = train_protein_ids+train_protein_ids[:args.batch_size]\ntest_protein_ids = test_protein_ids+test_protein_ids[:args.batch_size]\n\ntrain_protein_id_dataset = datasets.Dataset.from_dict({'protein_id': train_protein_ids})\ntest_protein_id_dataset = datasets.Dataset.from_dict({'protein_id': test_protein_ids})\n\ntrain_hf_dataset = train_protein_id_dataset.map(\n    process_protein_sequence,\n    input_columns='protein_id',\n    desc='Processing Protein Sequences',\n    num_proc=2 if args.debug else 16,\n).with_format('np')\n\ntest_hf_dataset = test_protein_id_dataset.map(\n    process_protein_sequence,\n    input_columns='protein_id',\n    desc='Processing Test Protein Sequences',\n    num_proc=2 if args.debug else 64,\n).with_format('np')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch \ndef collate_fn(batch_examples):\n    batch_input_ids = jnp.stack([jnp.array(example['input_ids'], dtype=jnp.int32) for example in batch_examples])\n    batch_attention_mask = jnp.stack([jnp.array(example['attention_mask'], dtype=jnp.int32) for example in batch_examples])\n    model_inputs = {\n        'input_ids': jnp.array(batch_input_ids, dtype=jnp.int32), \n        'attention_mask': jnp.array(batch_attention_mask, dtype=jnp.int32),\n    }\n    return model_inputs\n    \ntrain_dataloader = torch.utils.data.DataLoader(\n    train_hf_dataset,\n    batch_size=args.batch_size,\n    collate_fn=collate_fn,\n    drop_last=True, # DANGEROUS #\n)\ntest_dataloader = torch.utils.data.DataLoader(\n    test_hf_dataset,\n    batch_size=args.batch_size,\n    collate_fn=collate_fn,\n    drop_last=True, # DANGEROUS #\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model\n---\n","metadata":{}},{"cell_type":"code","source":"# General imports\nimport numpy as np\n\nfrom flax.core import frozen_dict\nimport jax\n\n# esmjax imports\nfrom esmjax import io, tokenizer as esm_tokenizer\nfrom esmjax.modules import modules\n\n# Imports specifically for multi-device sharding\nfrom esmjax.modules import partitioning\nfrom flax.linen import partitioning as nn_partitioning\nfrom jax.experimental import maps, PartitionSpec as P, pjit\n\nMODEL_NAME = \"esm2_t48_15B_UR50D\" # \"esm2_t6_8M_UR50D\" # \"esm2_t48_15B_UR50D\"  #\"esm2_t6_8M_UR50D\" #\"esm2_t36_3B_UR50D\" \n# Load in the original PyTorch state; will download if first time.\nstate = io.get_torch_state(MODEL_NAME)\n\nesm, params_axes = modules.get_esm2_model(state[\"cfg\"])\nesm_params = io.convert_encoder(state[\"model\"], state[\"cfg\"])\nesm_params = frozen_dict.FrozenDict({\"params\": esm_params})","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"state['cfg']","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"esm_params = flax.core.frozen_dict.unfreeze(esm_params)\ndel esm_params['params'][str(esm.num_layers-1)]\n# del params_axes[str(esm.num_layers-1)]\nesm.num_layers = esm.num_layers-1","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import jax.numpy as jnp\nimport flax.training.train_state\n\nTPU_MESH_SHAPE = (2, 4)\nTPU_SHARDING_RULES = [\n    (\"batch\", \"X\"),\n    (\"hidden\", \"Y\"),\n    (\"heads\", \"Y\"),\n    (\"embed_kernel\", \"X\"),\n    (\"embed\", \"Y\"),\n]\n\ndevices = np.asarray(jax.devices()).reshape(*TPU_MESH_SHAPE)\nglobal_mesh = jax.sharding.Mesh(devices=devices, axis_names=('X', 'Y'))\nesm_params = flax.core.frozen_dict.unfreeze(esm_params)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"esm_axes = partitioning.get_params_axes(esm_params, params_axes, rules=TPU_SHARDING_RULES)\npreshard_fn = pjit.pjit(\n    lambda x: x,  # this function does nothing\n    in_axis_resources=(esm_axes,),  # but this spec \"pre-shards\" the params\n    out_axis_resources=esm_axes,\n)\nwith global_mesh, nn_partitioning.axis_rules(TPU_SHARDING_RULES):\n    esm_params = flax.core.frozen_dict.freeze(esm_params)\n    esm_sharded_params = preshard_fn(esm_params)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create fn for inference.\n\n# 6:00 with batch size 16\n# 5:30 with batch size 64\n# 5:30 with batch size 256\nimport joblib\n\nesm_apply_fn = jax.experimental.pjit.pjit(\n    esm.apply,\n    in_shardings=(esm_axes, jax.experimental.PartitionSpec(\"X\", None)),\n    out_shardings=jax.experimental.PartitionSpec(\"X\", None, \"Y\"), # batch * seq * hidden\n)\n\ntrain_embeds = {'cls_embeds': [], 'sum_embeds': []}\nfor residue in IMP_RESIDUES:\n    train_embeds[residue] = []\n    \nwith global_mesh, nn_partitioning.axis_rules(TPU_SHARDING_RULES):\n    for step, batch in tqdm(enumerate(train_dataloader), total=len(train_dataloader)):\n        batch_start_idx, batch_end_idx = step*args.batch_size, (step+1)*args.batch_size\n        \n        batch_embeds = esm_apply_fn(esm_sharded_params, batch['input_ids'])\n        \n        def process_example(example_idx, example_sequence, example_embeds):\n            cls_embed = example_embeds[0, :]\n            sum_embed = np.sum(example_embeds, axis=0)\n            amino_acid_embeds = {acid: np.zeros(example_embeds.shape[-1], dtype=np.half) for acid in IMP_RESIDUES}\n            for seq_idx, amino_acid in enumerate(example_sequence.split()):\n                if seq_idx >= args.max_seq_len: \n                    continue\n                if amino_acid not in IMP_RESIDUES:\n                    continue\n                amino_acid_embeds[amino_acid] += example_embeds[seq_idx]\n            return cls_embed, sum_embed, amino_acid_embeds\n\n        # Parallelize the for loop using joblib\n        num_examples = batch_end_idx - batch_start_idx\n        \n        results = joblib.Parallel(n_jobs=4)(\n            joblib.delayed(process_example)(\n            example_idx, \n            protein_id_to_sequence[train_hf_dataset[example_idx]['protein_id']],\n            np.array(batch_embeds[example_idx - batch_start_idx], dtype=np.half),\n        ) for example_idx in range(batch_start_idx, batch_end_idx))\n        \n        # Extract the results from the parallel execution\n        cls_embeds, sum_embeds, amino_acid_embeds = zip(*results)\n        \n        # Store the results in the train_embeds dictionary\n        train_embeds['cls_embeds'].extend(cls_embeds)\n        train_embeds['sum_embeds'].extend(sum_embeds)\n        for ex_amino_acid_embeds in amino_acid_embeds:\n            for acid, v in ex_amino_acid_embeds.items():\n                train_embeds[acid].append(v)\n\n\nfor k, v in train_embeds.items():\n    filename = f'train_{k}.npy'\n    arr = np.stack(v)\n    np.save(filename, arr)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for k, v in tqdm(train_embeds.items()):\n    filename = f'train_{k}.npy'\n    arr = np.stack(v)\n    np.save(filename, arr)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 330 GB RAM for CPU\n# 1024 * 320","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}