{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":87793,"databundleVersionId":11228175,"sourceType":"competition"},{"sourceId":10967914,"sourceType":"datasetVersion","datasetId":6824140},{"sourceId":10983344,"sourceType":"datasetVersion","datasetId":6824149}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from datetime import datetime\nimport pytz\nprint('LOGGING TIME OF START:',  datetime.strftime(datetime.now(pytz.timezone('Asia/Singapore')), \"%Y-%m-%d %H:%M:%S\"))\n\n\ntry:\n\t#import zarr\n\tpass\nexcept:\n\tpass\n\nprint('PIP INSTALL OK !!!!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-10T15:14:37.492752Z","iopub.execute_input":"2025-03-10T15:14:37.493212Z","iopub.status.idle":"2025-03-10T15:14:37.532477Z","shell.execute_reply.started":"2025-03-10T15:14:37.493189Z","shell.execute_reply":"2025-03-10T15:14:37.531626Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"RIBONANZA_NET_DIR = '/kaggle/input/hengck23-rna-3d'\n\nimport sys\nsys.path.append(RIBONANZA_NET_DIR)\nsys.path.append(RIBONANZA_NET_DIR+'/ribonanza_net')\nfrom ribonanza_net.Network import *\nimport yaml\n\nimport pandas as pd\npd.set_option('display.max_columns', 20)\npd.set_option('display.expand_frame_repr', False)\n\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\nimport matplotlib \nimport matplotlib.pyplot as plt\n\n\n# helper--\nclass dotdict(dict):\n\t__setattr__ = dict.__setitem__\n\t__delattr__ = dict.__delitem__\n\n\tdef __getattr__(self, name):\n\t\ttry:\n\t\t\treturn self[name]\n\t\texcept KeyError:\n\t\t\traise AttributeError(name)\n\ndef set_aspect_equal(ax):\n\tx_limits = ax.get_xlim()\n\ty_limits = ax.get_ylim()\n\tz_limits = ax.get_zlim()\n\n\t# Compute the mean of each axis\n\tx_middle = np.mean(x_limits)\n\ty_middle = np.mean(y_limits)\n\tz_middle = np.mean(z_limits)\n\n\t# Compute the max range across all axes\n\tmax_range = max(x_limits[1] - x_limits[0],\n\t\t\t\t\ty_limits[1] - y_limits[0],\n\t\t\t\t\tz_limits[1] - z_limits[0]) / 2.0\n\n\t# Set the new limits to ensure equal scaling\n\tax.set_xlim(x_middle - max_range, x_middle + max_range)\n\tax.set_ylim(y_middle - max_range, y_middle + max_range)\n\tax.set_zlim(z_middle - max_range, z_middle + max_range)\n\n\nprint('torch',torch.__version__)\nprint('torch.cuda',torch.version.cuda)\n\nprint('IMPORT OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T15:14:37.533342Z","iopub.execute_input":"2025-03-10T15:14:37.533621Z","iopub.status.idle":"2025-03-10T15:14:42.822213Z","shell.execute_reply.started":"2025-03-10T15:14:37.533592Z","shell.execute_reply":"2025-03-10T15:14:42.821341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nMODE = 'submit' #'local' # submit\n\nDATA_KAGGLE_DIR = '/kaggle/input/stanford-rna-3d-folding'\nif MODE == 'local':\n\tfold0_index=[602, 603, 604, 605, 606, 607, 608, 609, 611, 612, 613, 614, 615, 616, 617, 618, 619, 620, 621, 622, 623, 624, 625, 627, 628, 629, 630, 631, 632, 633, 634, 635, 636, 637, 638, 640, 641, 642, 643, 644, 645, 646, 647, 648, 649, 650, 651, 652, 653, 654, 655, 656, 657, 658, 659, 660, 661, 662, 663, 664, 665, 666, 667, 668, 672, 673, 674, 677, 678, 679, 680, 681, 682, 683, 684, 685, 686, 688, 689, 690]\n\tvalid_df = pd.read_csv(f'{DATA_KAGGLE_DIR}/train_sequences.csv')\n\tvalid_df = valid_df.iloc[fold0_index].reset_index(drop=True)\n    \nif MODE == 'submit':\n\tvalid_df = pd.read_csv(f'{DATA_KAGGLE_DIR}/test_sequences.csv')\n\nprint('len(valid_df)',len(valid_df))\nprint(valid_df.iloc[0])\nprint('')\n\ncfg = dotdict(\n\tcheckpoint=\\\n    '/kaggle/input/hengck23-rna-3d-weight-00/00066666.pth',\n    #'/kaggle/input/hengck23-rna-3d-weight-00/00081300.pth',\n    #'/kaggle/input/hengck23-rna-3d-weight-00/lr-1e5-00113278.pth',\n    #'/kaggle/input/hengck23-rna-3d-weight-00/all_data-00062200.pth',\n    #'/kaggle/input/hengck23-rna-3d-weight-00/00108400.pth',\n\n    float_type=torch.float16, #32\n    max_length = 1024, #1600, #1024, #1344, #1400, #1280, #1024\n    dropout = 0.20,\n)\n\n\n\nprint('MODE:', MODE)\nprint('SETTING OK!!!')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T15:14:42.823059Z","iopub.execute_input":"2025-03-10T15:14:42.823420Z","iopub.status.idle":"2025-03-10T15:14:42.897725Z","shell.execute_reply.started":"2025-03-10T15:14:42.823385Z","shell.execute_reply":"2025-03-10T15:14:42.896855Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# model---\nclass Config:\n\tdef __init__(self, **entries):\n\t\tself.__dict__.update(entries)\n\t\tself.entries=entries\n\n\tdef print(self):\n\t\tprint(self.entries)\n\ndef load_config_from_yaml(file_path):\n\twith open(file_path, 'r') as file:\n\t\tconfig = yaml.safe_load(file)\n\treturn Config(**config)\n\n \nclass Net(RibonanzaNet):\n\tdef __init__(self,):\n\t\tconfig = load_config_from_yaml(f'{RIBONANZA_NET_DIR}/ribonanza_net/configs/pairwise.yaml')\n\t\tconfig.dropout = cfg.dropout\n        \n\t\tsuper(Net, self).__init__(config)\n\t\tself.D = nn.Parameter(torch.zeros(1))\n\t\tself.dropout = nn.Dropout(0.0)\n\t\tself.xyz_predictor = nn.Linear(256, 3)\n\n\tdef forward(self, token_id):\n\t\tdevice = self.D.device\n\t\ttoken_id = token_id.long().to(device)\n\t\ttoken_mask = torch.ones_like(token_id)\n\t\tB, L = token_id.shape\n\t\tsequence_feature, pairwise_feature = self.get_embeddings(token_id, token_mask)\n\t\txyz = self.xyz_predictor(sequence_feature)\n\t\treturn xyz\n\ntokenizer={nt:i for i,nt in enumerate('ACGU')}\nnet=Net()\n\n#print(net)\nprint(net.load_state_dict(\n\ttorch.load(\n\tcfg.checkpoint,\n\tmap_location=lambda storage, loc: storage, weights_only=False)['state_dict']\n))\n\n\nprint('MODEL OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T15:14:42.898663Z","iopub.execute_input":"2025-03-10T15:14:42.898879Z","iopub.status.idle":"2025-03-10T15:14:44.096722Z","shell.execute_reply.started":"2025-03-10T15:14:42.898860Z","shell.execute_reply":"2025-03-10T15:14:44.095876Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# start now !!!\n\ndef coord_to_df(sequence, coord, target_id):\n\tL = len(sequence)\n\tdf = pd.DataFrame()\n\tdf['ID'] = [f'{target_id}_{i + 1}' for i in range(L)]\n\tdf['resname'] = [s for s in sequence]\n\tdf['resid'] = [i + 1 for i in range(L)]\n\n\tnum_coord = len(coord)\n\tfor j in range(num_coord):\n\t\tdf[f'x_{j+1}'] = coord[j][:, 0]\n\t\tdf[f'y_{j+1}'] = coord[j][:, 1]\n\t\tdf[f'z_{j+1}'] = coord[j][:, 2]\n\treturn df\n\n\ndef run_submit(net):\n    net.cuda()\n \n    \n    submit_df=[]\n    for i,row in valid_df.iterrows():\n    \n        target_id = row.target_id\n        sequence = row.sequence\n        L = len(sequence)\n        coord = np.zeros((5,L,3),dtype=np.float32)\n        print('\\r', i, target_id, L, sequence[:10]+'...', end='', flush=True)\n        \n        seq = row.sequence[:cfg.max_length]  #<todo> handle long seq\n        token_id = [tokenizer[nt] for nt in seq]\n        token_id = torch.LongTensor(token_id).cuda()\n        \n    \n        #---------------------------------------------\n        # no dropout\n        j=0\n        net.eval()\n        \n        with torch.amp.autocast('cuda', dtype=cfg.float_type):\n            with torch.no_grad():\n                xyz = net(token_id.unsqueeze(0)).squeeze()\n                coord[j][:cfg.max_length] = xyz.cpu().data.numpy()\n        \n        #---------------------------------------------\n        # dropout\n        net.train()\n        for j in range(1,5):\n            with torch.amp.autocast('cuda', dtype=cfg.float_type):\n                with torch.no_grad():\n                    xyz = net(token_id.unsqueeze(0)).squeeze()\n                    coord[j][:cfg.max_length] = xyz.cpu().data.numpy()\n            \n        df = coord_to_df(sequence, coord, target_id)\n        submit_df.append(df)\n    \n        # print(df) \n        if i==0:\n            COLOR = ['red', 'blue', 'green', 'black', 'yellow', 'cyan', 'magenta']\n            fig = plt.figure(figsize=(10, 10))\n            ax = fig.add_subplot(111, projection='3d')\n            # ax.clear()\n            \n            for j in range(0, 5):\n                alpha =1 if j==0 else 0.3\n                x, y, z = coord[j][:, 0], coord[j][:, 1], coord[j][:, 2]\n                ax.scatter(x, y, z, c=COLOR[j], s=30, alpha=alpha)\n                ax.plot(x, y, z, color=COLOR[j], linewidth=1, alpha=alpha, label=f'{j}')\n            \n            set_aspect_equal(ax)\n            plt.legend()\n            plt.show()\n            # plt.waitforbuttonpress()\n            plt.close()\n        \n        #print(xyz)\n        #issue with long seq? will split join?\n    \n    torch.cuda.empty_cache()\n    print('')\n    \n    submit_df = pd.concat(submit_df)\n    submit_df.to_csv(f'submission.csv', index=False)\n    print(submit_df)\n    return submit_df\n\nrun_submit(net)\n\n\nprint('MODE:', MODE)\nprint('SUBMIT OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T15:14:44.097792Z","iopub.execute_input":"2025-03-10T15:14:44.098098Z","execution_failed":"2025-03-10T15:15:43.743Z"}},"outputs":[],"execution_count":null}]}