{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.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":12276181,"sourceType":"competition"},{"sourceId":13092586,"sourceType":"datasetVersion","datasetId":8292934},{"sourceId":13122157,"sourceType":"datasetVersion","datasetId":8293650},{"sourceId":311741,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":264400,"modelId":285488}],"dockerImageVersionId":31236,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport torch\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport torch\nimport random","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:43:04.293604Z","iopub.execute_input":"2026-01-12T00:43:04.293875Z","iopub.status.idle":"2026-01-12T00:43:08.849540Z","shell.execute_reply.started":"2026-01-12T00:43:04.293844Z","shell.execute_reply":"2026-01-12T00:43:08.848668Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference config","metadata":{"execution":{"iopub.status.busy":"2026-01-12T00:43:32.463841Z","iopub.execute_input":"2026-01-12T00:43:32.464622Z","iopub.status.idle":"2026-01-12T00:43:32.467667Z","shell.execute_reply.started":"2026-01-12T00:43:32.464593Z","shell.execute_reply":"2026-01-12T00:43:32.467075Z"}}},{"cell_type":"code","source":"out_folder=\"output_pdbs\"\n! mkdir output_pdbs\nnum_steps=200 #diffusion steps\nN_cycle=10 #number of recycling iters\nNpred=5 #number of preds, defaulted to top4 templates + 1 sequence only pred\nN_search=500 #number of sequences to search over for templates, sorted by pooled embedding sim\nN_search_ignored=500 ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:43:41.783547Z","iopub.execute_input":"2026-01-12T00:43:41.784018Z","iopub.status.idle":"2026-01-12T00:43:41.910417Z","shell.execute_reply.started":"2026-01-12T00:43:41.783991Z","shell.execute_reply":"2026-01-12T00:43:41.909654Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\n\nclass RNADataset(Dataset):\n    def __init__(self,data):\n        self.data=data\n        self.tokens={nt:i for i,nt in enumerate('ACGU')}\n\n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        sequence=[self.tokens[nt] for nt in self.data.loc[idx,'sequence']]\n        sequence=np.array(sequence)\n        sequence=torch.tensor(sequence)\n\n        return {'sequence':sequence}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:45:19.093199Z","iopub.execute_input":"2026-01-12T00:45:19.093961Z","iopub.status.idle":"2026-01-12T00:45:19.100322Z","shell.execute_reply.started":"2026-01-12T00:45:19.093926Z","shell.execute_reply":"2026-01-12T00:45:19.099503Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_data=pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\")\n\ntest_dataset=RNADataset(test_data)\ntest_dataset[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:46:12.061069Z","iopub.execute_input":"2026-01-12T00:46:12.061399Z","iopub.status.idle":"2026-01-12T00:46:12.133332Z","shell.execute_reply.started":"2026-01-12T00:46:12.061374Z","shell.execute_reply":"2026-01-12T00:46:12.132668Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\n\nsys.path.append(\"/kaggle/input/ribonanzanet2/pytorch/alpha/1\")\n\nfrom Network import *\nimport yaml\n\nclass Config:\n    def __init__(self, **entries):\n        self.__dict__.update(entries)\n        self.entries=entries\n\n    def print(self):\n        print(self.entries)\n\ndef load_config_from_yaml(file_path):\n    with open(file_path, 'r') as file:\n        config = yaml.safe_load(file)\n    return Config(**config)\n\nclass RibonanzaNet_embeddings(RibonanzaNet):\n    def __init__(self, config, pretrained=False):\n        config.dropout=0.1\n        config.use_grad_checkpoint=True\n        super(RibonanzaNet_embeddings, self).__init__(config)\n        if pretrained:\n            self.load_state_dict(torch.load(best_weights_path,map_location='cpu'))\n        self.dropout=nn.Dropout(0.0)\n\n\n\n    def forward(self,src):\n        \n        sequence_features, pairwise_features=self.get_embeddings(src, torch.ones_like(src).long().to(src.device))\n\n        # pairwise_features=pairwise_features+pairwise_features.permute(0,2,1,3)\n\n        # output=self.ct_predictor(self.dropout(pairwise_features))\n\n        return sequence_features, pairwise_features","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:46:21.180753Z","iopub.execute_input":"2026-01-12T00:46:21.181035Z","iopub.status.idle":"2026-01-12T00:46:24.377790Z","shell.execute_reply.started":"2026-01-12T00:46:21.181013Z","shell.execute_reply":"2026-01-12T00:46:24.377180Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from glob import glob\n\n\nmodel=RibonanzaNet_embeddings(load_config_from_yaml(\"/kaggle/input/ribonanzanet2/pytorch/alpha/1/pairwise.yaml\")).cuda()\nmodel.load_state_dict(torch.load(\"/kaggle/input/ribonanzanet2/pytorch/alpha/1/pytorch_model_fsdp.bin\",map_location='cpu'));\nmodel.eval();","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:46:27.637115Z","iopub.execute_input":"2026-01-12T00:46:27.637573Z","iopub.status.idle":"2026-01-12T00:46:34.135844Z","shell.execute_reply.started":"2026-01-12T00:46:27.637546Z","shell.execute_reply":"2026-01-12T00:46:34.135278Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\n\nembeddings=[]\n\nfor i in tqdm(range(len(test_dataset))):\n    example=test_dataset[i]\n    sequence=example['sequence'].cuda().unsqueeze(0)\n\n    with torch.no_grad():\n        sequence_features, pairwise_features=model(sequence)\n\n    embeddings.append(sequence_features.cpu().numpy())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:46:35.240704Z","iopub.execute_input":"2026-01-12T00:46:35.241235Z","iopub.status.idle":"2026-01-12T00:46:55.290872Z","shell.execute_reply.started":"2026-01-12T00:46:35.241193Z","shell.execute_reply":"2026-01-12T00:46:55.290239Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pickle\nwith open(\"/kaggle/input/rnet2-template-search-data/train_data.pkl\",\"rb\") as f:\n    data=pickle.load(f)\n\nwith open(\"/kaggle/input/rnet2-template-search-data/pdb_embeddings.p\",\"rb\") as f:\n    pdb_reacitivities=pickle.load(f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:48:28.992150Z","iopub.execute_input":"2026-01-12T00:48:28.992766Z","iopub.status.idle":"2026-01-12T00:49:46.307662Z","shell.execute_reply.started":"2026-01-12T00:48:28.992737Z","shell.execute_reply":"2026-01-12T00:49:46.307018Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pdb_reacitivities['CCUGGAUGGGA'].keys()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:49:46.309113Z","iopub.execute_input":"2026-01-12T00:49:46.309612Z","iopub.status.idle":"2026-01-12T00:49:46.314199Z","shell.execute_reply.started":"2026-01-12T00:49:46.309579Z","shell.execute_reply":"2026-01-12T00:49:46.313680Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data['ID']=[data['all_sequences'][i][:10] for i in range(len(data['all_sequences']))]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:49:46.315039Z","iopub.execute_input":"2026-01-12T00:49:46.315314Z","iopub.status.idle":"2026-01-12T00:49:46.328452Z","shell.execute_reply.started":"2026-01-12T00:49:46.315282Z","shell.execute_reply":"2026-01-12T00:49:46.327833Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sequences=data['sequence']\nc1_xyz=[]\nfor xyz in tqdm(data['xyz']):\n    c1=[]\n    for nt in xyz:\n        c1.append(nt[\"all\"][5])\n    c1=np.array(c1).astype('float32')\n    c1_xyz.append(c1)\nc1.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:49:46.329709Z","iopub.execute_input":"2026-01-12T00:49:46.329971Z","iopub.status.idle":"2026-01-12T00:49:49.749461Z","shell.execute_reply.started":"2026-01-12T00:49:46.329943Z","shell.execute_reply":"2026-01-12T00:49:49.748743Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cutoff_date=\"2025-05-29\"\ncutoff_date = pd.Timestamp(cutoff_date)\nalign_by=\"embedding\" \nprint(f\"Original number of sequences: {len(sequences)}\")\nfiltered_indices=[i for i in range(len(sequences)) if len(sequences[i])<=1024 and pd.Timestamp(data['temporal_cutoff'][i]) <= cutoff_date]\nsequences=[sequences[i] for i in filtered_indices]\nc1_xyz=[c1_xyz[i] for i in filtered_indices]\npdb_ids=[data['ID'][i] for i in filtered_indices]\nrelease_dates=[data['temporal_cutoff'][i] for i in filtered_indices]\nprint(f\"Number of sequences after filtering: {len(sequences)}\")\n\npdb_reacitivities=[pdb_reacitivities[sequence][align_by] for sequence in sequences]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:49:49.750473Z","iopub.execute_input":"2026-01-12T00:49:49.750778Z","iopub.status.idle":"2026-01-12T00:49:50.300189Z","shell.execute_reply.started":"2026-01-12T00:49:49.750746Z","shell.execute_reply":"2026-01-12T00:49:50.299145Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\ndef cosine_similarity(vec1, vec2):\n    \"\"\"\n    Compute cosine similarity between two vectors.\n    \n    Args:\n        vec1 (array-like): First vector\n        vec2 (array-like): Second vector\n    \n    Returns:\n        float: Cosine similarity in range [-1, 1]\n    \"\"\"\n    v1 = np.array(vec1)\n    v2 = np.array(vec2)\n    \n    dot = np.dot(v1, v2)\n    norm1 = np.linalg.norm(v1)\n    norm2 = np.linalg.norm(v2)\n    \n    if norm1 == 0 or norm2 == 0:\n        return 0.0  # define similarity as 0 if either vector is zero\n    \n    return dot / (norm1 * norm2)\n\ndef compute_similarity_matrix(A, B):\n    \"\"\"\n    Compute the similarity matrix SM (n x m) between embeddings A and B\n    using the exponential of the negative Euclidean distance.\n    \n    Args:\n        A: np.ndarray of shape (n, d)  -> embeddings of sequence A\n        B: np.ndarray of shape (m, d)  -> embeddings of sequence B\n        \n    Returns:\n        SM: np.ndarray of shape (n, m) -> similarity matrix\n    \"\"\"\n    # Pairwise squared Euclidean distances\n    distance = A[:, None, :] - B[None, :, :]   # (n, m, d)\n    dist2 = np.sum(distance ** 2, axis=-1)     # (n, m)\n\n    # Similarity = exp(-distance)\n    SM = np.exp(-np.sqrt(dist2 + 1e-8))        # add epsilon for stability\n    return SM\n\n\ndef enhance_similarity_matrix(SM):\n    \"\"\"\n    Enhance similarity matrix via Z-score normalization per row and column.\n    \n    Args:\n        SM: np.ndarray of shape (n, m)\n    \n    Returns:\n        SM_enh: np.ndarray of shape (n, m)\n    \"\"\"\n    # Row statistics\n    row_mean = SM.mean(axis=1, keepdims=True)   # (n, 1)\n    row_std  = SM.std(axis=1, keepdims=True) + 1e-8\n    z_row = (SM - row_mean) / row_std           # (n, m)\n\n    # Column statistics\n    col_mean = SM.mean(axis=0, keepdims=True)   # (1, m)\n    col_std  = SM.std(axis=0, keepdims=True) + 1e-8\n    z_col = (SM - col_mean) / col_std           # (n, m)\n\n    # Enhanced similarity\n    SM_enh = 0.5 * (z_row + z_col)\n    return SM_enh\n\n\ndef explicit_alignment_similarity(A, B):\n    \"\"\"\n    Full pipeline: compute similarity matrix and apply signal enhancement.\n    \n    Args:\n        A: np.ndarray of shape (n, d)\n        B: np.ndarray of shape (m, d)\n    \n    Returns:\n        SM: similarity matrix (n, m)\n        SM_enh: enhanced similarity matrix (n, m)\n    \"\"\"\n    SM = compute_similarity_matrix(A, B)\n    SM_enh = enhance_similarity_matrix(SM)\n    return SM_enh","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:49:50.301373Z","iopub.execute_input":"2026-01-12T00:49:50.301681Z","iopub.status.idle":"2026-01-12T00:49:50.311192Z","shell.execute_reply.started":"2026-01-12T00:49:50.301650Z","shell.execute_reply":"2026-01-12T00:49:50.310394Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from multiprocessing import Pool, cpu_count\n\n\n# Directions for traceback\nSTOP, DIAG, LEFT, UP = 0, 1, 2, 3\n\ndef generate_score_matrices(seq1, query_r, seq2, seq_r, match=3, mismatch=-3, gap_open=-10, gap_extend=-2):\n    \"\"\"\n    Smith-Waterman local alignment with affine gap penalty and traceback.\n    \"\"\"\n    n, m = len(seq1), len(seq2)\n\n    # DP matrices\n    H = np.zeros((n+1, m+1), dtype=float)\n    E = np.zeros((n+1, m+1), dtype=float)\n    F = np.zeros((n+1, m+1), dtype=float)\n\n    #E[:] = -1e9  # Initialize E to -inf to prevent gaps in seq1\n\n    # Traceback pointers: 0=STOP, 1=DIAG, 2=LEFT, 3=UP\n    traceback = np.zeros((n+1, m+1), dtype=int)\n\n    max_score = 0\n    max_pos = None\n    SM = explicit_alignment_similarity(query_r, seq_r) \n    for i in range(1, n+1):\n        for j in range(1, m+1):\n\n            # Match/mismatch\n            #score_sub = match if seq1[i-1] == seq2[j-1] else mismatch\n            #score_sub += 2*(0.5 - np.abs(query_r[i-1] - seq_r[j-1]).mean())*match\n            score_sub = SM[i-1, j-1]\n            #cos = cosine_similarity(query_r[i-1], seq_r[j-1])*match\n            #print(score_sub)\n            #score_sub = 10000\n            #print(reactivity_abs)\n            diag = H[i-1, j-1] + score_sub #+ cos\n\n            # Gaps (affine penalties)\n            E[i, j] = max(H[i, j-1] + gap_open, E[i, j-1] + gap_extend)\n            F[i, j] = max(H[i-1, j] + gap_open, F[i-1, j] + gap_extend)\n\n            # Best score at this cell\n            H[i, j] = max(0, diag, E[i, j], F[i, j])\n            \n            action = np.argmax([0, diag, E[i, j], F[i, j]])\n\n            # Traceback direction\n            if action == 0:\n                traceback[i, j] = STOP\n            elif action == 1:\n                traceback[i, j] = DIAG\n            elif action == 2:\n                traceback[i, j] = LEFT\n            elif action == 3:\n                traceback[i, j] = UP\n\n            if H[i, j] > max_score:\n                max_score = H[i, j]\n                max_pos = (i, j)\n\n    return H, max_score, max_pos, traceback\n\n\n\n\ndef traceback_alignment(seq1, seq2, xyz, traceback, start_pos):\n    \"\"\" Reconstruct the alignment from traceback matrix. \"\"\"\n    aligned1, aligned2, template = [], [], []\n    #print(start_pos)\n    i, j = start_pos\n\n    if i < len(seq1)+1: # add seq1 unaligned prefix\n        for k in range(len(seq1),i,-1):\n            aligned1.append(seq1[k-1])\n            aligned2.append('-')\n            template.append(np.array([np.nan, np.nan, np.nan]))  # No coordinates for gaps\n\n    while traceback[i, j] != STOP:\n        if traceback[i, j] == DIAG:\n            if seq1[i-1]==seq2[j-1]:\n                \n                aligned1.append(seq1[i-1])\n                aligned2.append(seq2[j-1])\n            else: #lowercase for mismatches\n                aligned1.append(seq1[i-1].lower())\n                aligned2.append(seq2[j-1].lower())\n            template.append(xyz[j-1])\n            i, j = i-1, j-1\n        elif traceback[i, j] == LEFT:\n            aligned1.append('-')\n            aligned2.append(seq2[j-1])\n            j -= 1\n        elif traceback[i, j] == UP:\n            aligned1.append(seq1[i-1])\n            aligned2.append('-')\n            template.append(np.array([np.nan, np.nan, np.nan]))\n            i -= 1\n\n    while i > 0:\n        aligned1.append(seq1[i-1])\n        aligned2.append('-')\n        template.append(np.array([np.nan, np.nan, np.nan]))\n        #aligned2.append(seq2[j-1])\n        i -= 1\n\n    template = np.array(template)[::-1]  # Reverse to correct order\n\n    #alignment = ''.join(reversed(aligned2))\n\n    return ''.join(reversed(aligned1)), ''.join(reversed(aligned2)), template\n\n\n    #return ''.join(reversed(aligned1)), ''.join(reversed(aligned2))\n\ndef smith_waterman_affine_gap(seq1, query_r, seq2, seq_r, xyz, pdb_id, match=3, mismatch=-3, gap_open=-10, gap_extend=-2):\n    H, score, pos, tb = generate_score_matrices(seq1, query_r, seq2, seq_r, match, mismatch, gap_open, gap_extend)\n    alignment1,  alignment2, template = traceback_alignment(seq1, seq2, xyz, tb, pos)\n\n    score = score/ max(len(seq1), len(seq2))  # normalize by length of longer sequence\n    return {\"alignment1\":alignment1,\n            \"alignment2\":alignment2,\n            \"template\":template,\n            \"score\":score,\n            \"query\":seq1,\n            \"target\":seq2,\n            \"pdb_id\":pdb_id}\n\n# helper wrapper so Pool.map can pickle it\ndef _smith_wrapper(args):\n    query, query_r, seq, seq_r, xyz, id = args\n\n    result = smith_waterman_affine_gap(query, query_r, seq, seq_r, xyz, id)\n    return result\npool = Pool(processes=cpu_count())\ndef get_template(query, query_r, sequences, pdb_reacitivities, c1_xyz, ID, PDB_IDs, processes=None, top_search=500, topk=5):\n    results = []\n\n    #compute cosine between query_r and all pdb_reactivities\n    #then only take top 100 in sequences, c1_xyz, pdb_reacitivities,PDB_IDs\n    #use multi-processing to compute cosine\n    print(\"Computing cosine similarities...\")\n    cosines=[cosine_similarity(query_r.mean(0), pdb_reacitivities[j].mean(0)) for j in tqdm(range(len(pdb_reacitivities)))]\n    l2_sim=[-np.linalg.norm(query_r.mean(0)-pdb_reacitivities[j].mean(0)) for j in tqdm(range(len(pdb_reacitivities)))]\n    cosines=np.array(cosines)\n    l2_sim=np.array(l2_sim)\n    #take z scores and sum\n    cosines=(cosines - np.mean(cosines)) / np.std(cosines)\n    l2_sim=(l2_sim - np.mean(l2_sim)) / np.std(l2_sim)\n    sim=0.5*cosines+0.5*l2_sim\n\n    #pool = Pool(processes=cpu_count())\n    # with Pool(processes=nproc) as pool:\n    #     cosines = []\n    #     for cos in tqdm(pool.imap_unordered(lambda r: cosine_similarity(query_r.mean(0), r.mean(0)), pdb_reacitivities), total=len(pdb_reacitivities)):\n    #         cosines.append(cos)\n    top_indices = np.argsort(sim)[-top_search:][::-1]\n    sequences=[sequences[i] for i in top_indices]\n    c1_xyz=[c1_xyz[i] for i in top_indices]\n    pdb_reacitivities=[pdb_reacitivities[i] for i in top_indices]\n    PDB_IDs=[PDB_IDs[i] for i in top_indices]\n\n\n    # prepare argument tuples\n    args_list = [(query, query_r, seq, seq_r, xyz, pdb_id) for seq, xyz, seq_r, pdb_id in zip(sequences, c1_xyz, pdb_reacitivities,PDB_IDs)]\n\n    # decide number of processes\n    nproc = processes or cpu_count()\n\n    #$with Pool(processes=nproc) as pool:\n        # imap_unordered gives streaming results + tqdm progress bar\n    for r in tqdm(pool.imap_unordered(_smith_wrapper, args_list), total=len(args_list)):\n        results.append(r)\n    # for args in tqdm(args_list):\n    #     r=_smith_wrapper(args)\n    #     results.append(r)\n\n    # rank top-k\n    best_indices = np.argsort([r['score'] for r in results])[-topk:][::-1]\n    \n    # gather top-k results and put into DataFrame\n    alignments = [results[i] for i in best_indices[:topk]]\n    alignments = pd.DataFrame(alignments)\n\n    # build output DataFrame\n    rows = []\n    \n    for i in range(len(query)):\n        row = [f\"{ID}_{i+1}\", query[i], i+1]\n        for j in range(5):\n            row.extend(list(results[best_indices[j]]['template'][i]))\n        rows.append(row)\n\n    df = pd.DataFrame(rows, columns=example_template_file.columns)\n    return df, alignments","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:49:50.312320Z","iopub.execute_input":"2026-01-12T00:49:50.312649Z","iopub.status.idle":"2026-01-12T00:49:50.850339Z","shell.execute_reply.started":"2026-01-12T00:49:50.312614Z","shell.execute_reply":"2026-01-12T00:49:50.848863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"example_template_file=pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/sample_submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:49:55.786468Z","iopub.execute_input":"2026-01-12T00:49:55.786831Z","iopub.status.idle":"2026-01-12T00:49:55.821100Z","shell.execute_reply.started":"2026-01-12T00:49:55.786789Z","shell.execute_reply":"2026-01-12T00:49:55.820282Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\n\nnot_ignored={'16a3', '6c24', '1deb', 'bc4e', '9886', 'New5', '027a', '6406','New4', 'ce51'}\n\nstart= time.time()\ndfs=[]\nalignments=[]\nfor i in range(len(test_data)):\n#for i in range(10):\n    query=test_data.loc[i,'sequence']\n    ID=test_data.loc[i,'target_id']\n    query_r=embeddings[i].squeeze()\n    if ID in not_ignored:\n        top_search=N_search\n    else:\n        top_search=N_search_ignored #only search top10 based on pooled embedding similarity for ignored targets to save time\n    template,alignment = get_template(query,query_r,sequences,pdb_reacitivities,c1_xyz,ID,pdb_ids,top_search=top_search,topk=5)\n    dfs.append(template)\n    alignment['target_id'] = ID\n    alignments.append(alignment)\n    #break\ndfs=pd.concat(dfs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-12T00:50:59.887626Z","iopub.execute_input":"2026-01-12T00:50:59.888361Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"template_csv=dfs\ntemplate_csv.to_csv('templates.csv',index=False)\ntemplate_csv.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Rnet3D w templates\n","metadata":{}},{"cell_type":"code","source":"def get_template(target_id, idx=1):\n\n    if templates is None:\n        return None\n\n    rows=templates[templates['target_id']==target_id].reset_index(drop=True)\n\n    xyz = [rows[f'x_{idx}'].values,rows[f'y_{idx}'].values,rows[f'z_{idx}'].values]\n    xyz = np.array(xyz).T\n    \n    xyz[xyz<-1e9]=0.0\n    xyz[xyz!=xyz]=0.0\n\n    nan_fraction = (xyz == 0.0).mean()\n\n    if nan_fraction < 0.8:\n        print(f\"{target_id} has template w {nan_fraction} missing\")\n        xyz[xyz==0.0]=np.nan\n        xyz=torch.tensor(xyz)\n        distance_matrix_input = calculate_distance_matrix(xyz, xyz)\n        distance_matrix_input[torch.isnan(distance_matrix_input)] = 0.0\n        distance_matrix_input = distance_matrix_input.clip(0, 39).long().cuda()\n\n        return distance_matrix_input\n    else:\n        print(f\"{target_id} has no template w {nan_fraction} missing\")\n        return None","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open(\"/kaggle/input/rnet3d-test141/nt_reference_positions.pkl\", \"rb\") as f:\n    reference_positions = pickle.load(f)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\n#sys.path.append(\"/kaggle/input/rnet3d-test141\") \n\n!cp /kaggle/input/rnet3d-test141/*py .\n\nimport pandas as pd\nimport torch\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport torch\nimport random\nimport pickle\nfrom utils import *\nimport os\n#from Diffusion import Diffusion\nimport argparse\nfrom ALL_ATOM import *\nfrom collections import defaultdict\n\n#count for each group\nrna_atom_group_counts = {k: len(v[\"all\"]) for k, v in rna_atom_groups.items()}\nprint(\"RNA Atom Group Counts:\", rna_atom_group_counts)\n\nclass RNA3D_TestDataset(Dataset):\n    def __init__(self,data):\n        self.data=data\n        #set default to 4\n        self.tokens=defaultdict(lambda: 4)\n        self.tokens['A']=0\n        self.tokens['C']=1\n        self.tokens['G']=2\n        self.tokens['U']=3\n\n        self.data=data\n\n\n        #{nt:i for i,nt in enumerate('ACGU')}\n\n    def __len__(self):\n        return len(self.data)\n    \n    def get_all_atom_data(self, sequence, res_ids):\n        all_atom_index=[]\n        all_atom_res_index=[]\n        all_atom_ref_positions=[]\n        for i in range(len(sequence)):\n            nt=sequence[i]\n\n            nt_str = 'ACGU'[nt]\n\n            assert nt_str in 'ACGU'\n\n            offset = 0\n            ref_pos=random_rotation_point_cloud_torch(torch.tensor(reference_positions[nt_str]))\n            all_atom_ref_positions.append(ref_pos)\n\n            n_atoms = len(ref_pos)\n\n            atom_names = rna_atom_groups[nt_str]['all']\n            \n            #get atom names and types and then convert to one hot indices\n            atom_indices = np.array([atom_name_one_hots[name] for name in atom_names])\n            atom_types = np.array([atom_type_one_hots[name[0]] for name in atom_names])\n\n            atom_features = np.concatenate([atom_indices, atom_types], axis=-1)\n            # print(f\"Atom features shape: {atom_features.shape}\")\n            # exit()\n\n\n\n\n            #all_atom_index.append(nt*30 + np.arange(n_atoms))\n            all_atom_index.append(atom_features)\n            all_atom_res_index.append([i]*n_atoms)\n\n            \n\n\n            #print(len(res_xyz))\n\n        all_atom_index=np.concatenate(all_atom_index)\n        all_atom_res_index=np.concatenate(all_atom_res_index)   \n        all_atom_ref_positions=torch.cat(all_atom_ref_positions)#/config.data_std\n\n        #convert to torch tensors\n        all_atom_index=torch.tensor(all_atom_index, dtype=torch.float32)\n        all_atom_res_index=torch.tensor(all_atom_res_index, dtype=torch.long)\n        #all_atom_ref_positions=torch.tensor(all_atom_ref_positions, dtype=torch.float32)\n\n        all_atom_index = torch.cat([all_atom_ref_positions,all_atom_index],-1)\n\n        return all_atom_index, all_atom_res_index\n\n    def __getitem__(self, idx):\n\n        sequence=[self.tokens[nt] for nt in (self.data.loc[idx,'sequence'])]\n        sequence=np.array(sequence)\n        res_ids=np.arange(len(sequence))\n\n        all_atom_index, all_atom_res_index  = self.get_all_atom_data(sequence, res_ids)       \n\n        sequence=torch.tensor(sequence, dtype=torch.long)\n        res_ids=torch.tensor(res_ids, dtype=torch.long)\n\n        return {'sequence':sequence,\n                'res_ids':res_ids,\n                'all_atom_index':all_atom_index,\n                'all_atom_res_index':all_atom_res_index,}\n\ntest_dataset=RNA3D_TestDataset(test_data)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import importlib.util\nimport sys\n\nspec = importlib.util.spec_from_file_location(\n    \"Network\", \"/kaggle/input/rnet3d-test141/Network.py\"\n)\nNetwork = importlib.util.module_from_spec(spec)\nspec.loader.exec_module(Network)\n\n# Equivalent of `from Network import *`\nglobals().update({k: v for k, v in Network.__dict__.items() if not k.startswith(\"_\")})","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#model definition\nimport torch.nn as nn\nimport torch.nn as nn\nimport torch\nfrom utils import random_rotation_point_cloud_torch_batch\nfrom Network import *\nfrom all_atom_module import AllAtomEncoder, AllAtomDecoder\nfrom TemplateEmbedder import TemplateEmbedder\n\nimport torch\nimport torch.nn as nn\nimport math\n\nimport torch\nimport torch.nn as nn\nimport math\n\n\nclass TemplateEmbedder(nn.Module):\n    def __init__(self, ninp=384, pairwise_dim=128, nhead=12, nlayer=4, dim_msa=32):\n        super(TemplateEmbedder, self).__init__()\n        self.distance_matrix_embedder = nn.Embedding(45, pairwise_dim)\n\n        self.encoder = nn.Embedding(100, ninp)\n        self.outer_product_mean=Outer_Product_Mean(in_dim=ninp,dim_msa=dim_msa,pairwise_dim=pairwise_dim)\n        self.pos_encoder=relpos(pairwise_dim)\n\n        self.blocks = []\n\n        for i in range(nlayer):\n            layer=FoldingBlock(ninp, \n            nhead, \n            ninp*4, \n            pairwise_dim,\n            False, 32)\n            self.blocks.append(layer)\n\n        self.blocks=nn.ModuleList(self.blocks)\n\n        self.post_norm = nn.LayerNorm(pairwise_dim)\n        #self.valid_distance_embedder = nn.Linear(1, pairwise_dim, bias=False)\n\n    def forward(self, sequence, distance_matrix_input):\n\n        mask = torch.ones_like(sequence)\n        sequence_features = self.encoder(sequence)\n        distance_matrix_input[torch.isnan(distance_matrix_input)] = 0.0\n        distance_matrix_input = self.distance_matrix_embedder(distance_matrix_input)\n        \n        \n        pairwise_features = self.outer_product_mean(sequence_features)\n        pairwise_features = self.pos_encoder(pairwise_features)\n\n        pairwise_features = pairwise_features + distance_matrix_input.unsqueeze(0)\n        \n        for block in self.blocks:\n            sequence_features,pairwise_features=checkpoint.checkpoint(block, \n            [sequence_features, pairwise_features, mask , False],\n            use_reentrant=False)\n\n        pairwise_features = self.post_norm(pairwise_features)\n        return sequence_features, pairwise_features\n\nclass FourierEmbedding(nn.Module):\n    def __init__(self, embed_dim, sigma_data, n_freq = 256):\n        \"\"\"\n        Args:\n            embed_dim (int): Dimensionality of the Fourier embedding\n            sigma_data (float): The σ̂_data constant for log transformation\n        \"\"\"\n        super().__init__()\n        self.embed_dim = embed_dim\n        self.sigma_data = sigma_data\n\n        # Fixed random weights and biases for Fourier features\n        self.register_buffer('w', torch.randn(n_freq))\n        self.register_buffer('b', torch.randn(n_freq))\n\n        # LayerNorm without affine parameters\n        self.layernorm = nn.LayerNorm(n_freq, elementwise_affine=True)\n\n        # Linear projection without bias\n        self.linear = nn.Linear(n_freq, embed_dim, bias=False)\n\n    def forward(self, t):\n        \"\"\"\n        Args:\n            t (Tensor): A 1D or 2D tensor of shape [batch] or [batch, 1]\n        Returns:\n            Tensor: Transformed time embedding, shape [batch, embed_dim]\n        \"\"\"\n        if t.dim() == 1:\n            t = t[:, None]  # shape [batch, 1]\n        t_input = 0.25 * torch.log(t / self.sigma_data)  # shape [batch, 1]\n\n        x = 2 * math.pi * (t_input @ self.w[None, :] + self.b)  # [batch, embed_dim]\n        x = torch.cos(x)\n        x = self.layernorm(x)\n        x = self.linear(x)\n        return x\n\n\n\nclass finetuned_RibonanzaNet(RibonanzaNet):\n    def __init__(self, rnet_config, config, pretrained=False):\n        rnet_config.dropout=config.trunk_dropout\n        rnet_config.use_grad_checkpoint=True\n        self.rnet_config=rnet_config\n        super(finetuned_RibonanzaNet, self).__init__(rnet_config)\n        if pretrained:\n            self.load_state_dict(torch.load(config.pretrained_weight_path,map_location='cpu'))\n        # self.ct_predictor=nn.Sequential(nn.Linear(64,256),\n        #                                 nn.ReLU(),\n        #                                 nn.Linear(256,64),\n        #                                 nn.ReLU(),\n        #                                 nn.Linear(64,1)) \n        self.dropout=nn.Dropout(0.0)\n\n        #rnet_config.nlayers=48\n        rnet_config.dropout=0.1\n        self.extra_evoformer = []\n\n        for i,layer in enumerate(range(config.folding_blocks)):\n\n            layer=Network.FoldingBlock(rnet_config.ninp, \n            rnet_config.nhead, \n            rnet_config.ninp*4, \n            rnet_config.pairwise_dimension,\n            False, rnet_config.dim_msa)\n\n            #scale_factor=1/(i+1)**0.5\n            #scale_factor=i+1\n            #scale_factor=0\n            #recursive_linear_init(layer,scale_factor)\n            #zero_init(layer.attn_linear)\n            # zero_init(layer.sequence_transititon[2])\n            # zero_init(layer.pair_transition[3])\n            # zero_init(layer.triangle_update_out.to_out)\n            # zero_init(layer.triangle_update_in.to_out)\n\n            self.extra_evoformer.append(layer)\n        self.extra_evoformer=nn.ModuleList(self.extra_evoformer)\n\n        self.recycle_sequence=nn.Sequential(nn.LayerNorm(rnet_config.ninp),\n                                            nn.Linear(rnet_config.ninp,rnet_config.ninp,bias=False))\n        self.recycle_pairwise=nn.Sequential(nn.LayerNorm(rnet_config.pairwise_dimension),\n                                            nn.Linear(rnet_config.pairwise_dimension,rnet_config.pairwise_dimension,bias=False))\n\n        zero_init(self.recycle_sequence[1])\n        zero_init(self.recycle_pairwise[1])\n\n        #self.layer_weights=torch.zeros(48, dtype=torch.float32)\n        self.layer_weights=torch.linspace(0, 1, 48, dtype=torch.float32)\n        self.layer_weights[-1]=-1e18\n        self.layer_weights=nn.Parameter(self.layer_weights,requires_grad=True)\n\n\n        decoder_dim=config.decoder_dim\n        self.structure_module=[SimpleStructureModule(d_model=decoder_dim, nhead=config.decoder_nhead, \n                 s_model=rnet_config.ninp,\n                 dim_feedforward=decoder_dim*2, pairwise_dimension=rnet_config.pairwise_dimension, dropout=0.0) for i in range(config.decoder_num_layers)]\n        self.structure_module=nn.ModuleList(self.structure_module)\n\n        for i,layer in enumerate(self.structure_module):\n            scale_factor=1/(i+1)**0.5\n            #scale_factor=i+1\n            #scale_factor=0\n            #recursive_linear_init(layer,scale_factor)\n            #zero_init(layer.attn_linear)\n            #zero_init(layer.conditioned_transiton.output_proj)\n\n\n        self.xyz_embedder=nn.Linear(3,rnet_config.ninp)\n        self.xyz_norm=nn.LayerNorm(rnet_config.ninp)\n        self.xyz_predictor=nn.Sequential(nn.LayerNorm(decoder_dim),\n                                         nn.Linear(decoder_dim,3))\n                                            \n        \n\n\n        self.distogram_predictor=nn.Sequential(nn.LayerNorm(rnet_config.pairwise_dimension),\n                                                nn.Linear(rnet_config.pairwise_dimension,40))\n\n        self.time_embedder=FourierEmbedding(rnet_config.ninp, sigma_data=config.data_std)\n\n        self.concat_linear=nn.Sequential(nn.Linear(rnet_config.ninp+decoder_dim+3,decoder_dim),nn.LayerNorm(decoder_dim))\n\n        self.time_mlp1=TransitionLayer(rnet_config.ninp,2)\n        self.time_mlp2=TransitionLayer(rnet_config.ninp,2)\n        \n\n        self.time_mlp3=nn.Sequential(\n                                     nn.Linear(rnet_config.ninp,rnet_config.ninp*2),\n                                     nn.ReLU(),  \n                                     nn.Linear(rnet_config.ninp*2,rnet_config.ninp*1))\n        self.time_norm3=nn.LayerNorm(rnet_config.ninp)\n\n        self.tgt_norm=nn.LayerNorm(rnet_config.ninp)\n\n        self.distance2pairwise=nn.Linear(1,rnet_config.pairwise_dimension,bias=False)\n\n        self.pair_mlp1=TransitionLayer(rnet_config.pairwise_dimension,2)\n\n        self.pair_mlp2=TransitionLayer(rnet_config.pairwise_dimension,2)\n\n        self.pair_distance_mlp=nn.Sequential(nn.Linear(4,8),\n                                                nn.ReLU(),\n                                                nn.Linear(8,4))\n        self.pair_distance_mlp2=nn.Sequential(  nn.LayerNorm(4),\n                                                nn.Linear(4,8),\n                                                nn.ReLU(),\n                                                nn.Linear(8,4))\n\n        self.pair_vector_linear=nn.Linear(3,rnet_config.pairwise_dimension,bias=False)\n\n        #hyperparameters for diffusion\n        self.n_times = config.n_times\n\n        #self.model = model\n        \n        # define linear variance schedule(betas)\n        beta_1, beta_T = config.beta_min, config.beta_max\n        betas = torch.linspace(start=beta_1, end=beta_T, steps=config.n_times)#.to(device) # follows DDPM paper\n        self.sqrt_betas = torch.sqrt(betas)\n                                     \n        # define alpha for forward diffusion kernel\n        self.alphas = 1 - betas\n        self.sqrt_alphas = torch.sqrt(self.alphas)\n        alpha_bars = torch.cumprod(self.alphas, dim=0)\n        self.sqrt_one_minus_alpha_bars = torch.sqrt(1-alpha_bars)\n        self.sqrt_alpha_bars = torch.sqrt(alpha_bars)\n\n        self.data_std=config.data_std\n\n        self.adaptor=nn.Linear(rnet_config.ninp,config.decoder_dim,bias=False)\n\n        self.all_atom_encoder = AllAtomEncoder(max_len=config.max_len)\n\n        self.all_atom_avg = nn.Linear(128, 128)\n        self.all_atom_upsample = nn.Linear(128, config.decoder_dim)\n\n        self.all_atom_decoder = AllAtomDecoder(max_len=config.max_len)\n\n        self.tgt_downsample = nn.Sequential(nn.LayerNorm(config.decoder_dim),\n                                            nn.Linear(config.decoder_dim, 128))\n        self.s2a = nn.Sequential(nn.LayerNorm(rnet_config.ninp),\n                                    nn.Linear(rnet_config.ninp, config.decoder_dim))\n\n        self.template_embedder = TemplateEmbedder()\n\n        # self.sequence_features_norm = nn.LayerNorm(rnet_config.ninp)\n\n        self.s_conditioning_linear=nn.Sequential(nn.LayerNorm(rnet_config.ninp),nn.Linear(rnet_config.ninp,rnet_config.ninp,bias=False))\n\n    def custom(self, module):\n        def custom_forward(*inputs):\n            inputs = module(*inputs)\n            return inputs\n        return custom_forward\n    \n    def embed_pair_distance(self,inputs):\n        pairwise_features,xyz=inputs\n        vector_matrix=(xyz[:,None,:,:]-xyz[:,:,None,:])#*self.data_std\n\n        distance_matrix=(vector_matrix**2).sum(-1)\n        distance_matrix=1/(1+distance_matrix)\n        distance_matrix=distance_matrix[:,:,:,None]\n        # pairwise_features=pairwise_features+\\\n        #                   self.distance2pairwise(distance_matrix)+\\\n        #                   self.pair_vector_linear(vector_matrix)\n\n        # pairwise_features+=self.pair_mlp1(pairwise_features)\n        # pairwise_features+=self.pair_mlp2(pairwise_features)\n        distance_features=torch.cat([vector_matrix,distance_matrix],-1)\n        distance_features=distance_features+self.pair_distance_mlp(distance_features)\n        distance_features=distance_features+self.pair_distance_mlp2(distance_features)\n        return distance_features\n\n    def get_embeddings(self, src,src_mask=None,return_aw=False,res_ids=None):\n        B,L=src.shape\n        src = src\n        src = self.encoder(src).reshape(B,L,-1)\n        \n        #spawn outer product\n        if self.use_gradient_checkpoint:\n            #print(\"using grad checkpointing\")\n            pairwise_features=checkpoint.checkpoint(self.custom(self.outer_product_mean), src, use_reentrant=False)\n            pairwise_features=pairwise_features+self.pos_encoder(src)\n        else:\n            pairwise_features=self.outer_product_mean(src)\n            pairwise_features=pairwise_features+self.pos_encoder(src)\n\n\n        #attention_weights=[]\n\n        if self.training:\n            all_sequence_features=[]\n            all_pairwise_features=[]\n            for i,layer in enumerate(self.transformer_encoder):\n                src,pairwise_features=checkpoint.checkpoint(self.custom(layer), \n                [src, pairwise_features, src_mask, return_aw],\n                use_reentrant=False)\n\n                all_sequence_features.append(src)\n                all_pairwise_features.append(pairwise_features)\n\n            return torch.stack(all_sequence_features,0), torch.stack(all_pairwise_features,0)\n        else:\n            all_sequence_features=torch.zeros_like(src)\n            all_pairwise_features=torch.zeros_like(pairwise_features)\n            layer_weights=self.layer_weights.softmax(0)#[:,None,None,None]\n            for i,layer in enumerate(self.transformer_encoder):\n                src,pairwise_features=layer([src, pairwise_features, src_mask, return_aw])\n\n\n                all_sequence_features+=src * layer_weights[i,None,None,None]\n                all_pairwise_features+=pairwise_features * layer_weights[i,None,None,None,None]\n\n            return all_sequence_features, all_pairwise_features\n\n\n    def get_conditioning(self,src,distance_matrix_input,cycles,trunk_grad=False,res_ids=None):\n        \n        #print(f'Get conditioning with trunk_grad={trunk_grad} and cycles={cycles}')\n        with torch.set_grad_enabled(trunk_grad):\n            all_sequence_features, all_pairwise_features=self.get_embeddings(src, torch.ones_like(src).long().to(src.device),\n            res_ids=res_ids)\n            #print(f'Got embeddings with shape {all_sequence_features.shape} and {all_pairwise_features.shape}')\n            #exit()\n\n        B,L = src.shape\n        s_hat=torch.zeros(B,L,self.rnet_config.ninp).to(src.device)\n        z_hat=torch.zeros(B,L,L,self.rnet_config.pairwise_dimension).to(src.device)\n\n        for c in range(1,cycles+1):\n            #print(f'Cycle {c}/{cycles}')\n            with torch.set_grad_enabled(self.training and c==cycles):\n            #with torch.no_grad():\n                if self.training:\n                    sequence_features=all_sequence_features*self.layer_weights.softmax(0)[:,None,None,None]\n                    pairwise_features=all_pairwise_features*self.layer_weights.softmax(0)[:,None,None,None,None]\n                    # print(self.layer_weights.softmax(0))\n                    # exit()\n                    sequence_features_init=sequence_features.sum(0)\n                    pairwise_features_init=pairwise_features.sum(0)\n                else:\n                    sequence_features_init=all_sequence_features\n                    pairwise_features_init=all_pairwise_features\n\n                sequence_features=sequence_features_init+self.recycle_sequence(s_hat)\n                pairwise_features=pairwise_features_init+self.recycle_pairwise(z_hat)\n\n                mask = torch.ones_like(src).long().to(src.device)\n                if c==cycles:\n                    self.extra_evoformer.requires_grad_(True)\n                    self.template_embedder.requires_grad_(True)\n                else:\n                    self.extra_evoformer.requires_grad_(False)\n                    self.template_embedder.requires_grad_(False)\n\n                #print(f'running cycle {c} with trunk_grad={trunk_grad} and cycles={cycles}')\n                if distance_matrix_input is not None:\n                    template_s, template_z = self.template_embedder(src, distance_matrix_input)\n                    sequence_features = sequence_features + template_s\n                    pairwise_features = pairwise_features + template_z\n\n\n\n                for layer in self.extra_evoformer:\n                    #with torch.autocast(device_type=\"cuda\", dtype=torch.bfloat16):\n                    if c==cycles:\n                        sequence_features,pairwise_features=checkpoint.checkpoint(layer, \n                        [sequence_features, pairwise_features, mask , False],\n                        use_reentrant=False)\n                    else:\n                        sequence_features,pairwise_features=layer([sequence_features, pairwise_features, mask.to(src.device), False])\n                    #sequence_features,pairwise_features=layer([sequence_features, pairwise_features, torch.ones_like(src).long().to(src.device), False])\n                s_hat=sequence_features\n                z_hat=pairwise_features\n\n                # if c!=cycles:\n                #     s_hat=s_hat.detach()\n                #     z_hat=z_hat.detach()\n\n        return sequence_features,pairwise_features\n\n    def get_decoder_features(self, sequence_features, pairwise_features, t):\n        decoder_batch_size=len(t)\n        # print(t.shape)\n        # exit()\n        sequence_features=self.s_conditioning_linear(sequence_features)\n        sequence_features=sequence_features.repeat(decoder_batch_size,1,1)\n        \n\n\n        time_embed=self.time_embedder(t).unsqueeze(1)\n        # print(f\"Time embedding shape: {time_embed.shape}\")\n        # exit()\n        #xyz=self.xyz_norm(self.xyz_embedder(xyz))\n        \n\n\n        tgt=time_embed+sequence_features\n        tgt=tgt + self.time_mlp1(tgt)\n        tgt=tgt + self.time_mlp2(tgt)\n\n        pairwise_features = pairwise_features+self.pair_mlp1(pairwise_features)\n        pairwise_features = pairwise_features+self.pair_mlp2(pairwise_features)\n        \n        return tgt, pairwise_features\n\n    def forward(self,src,distance_matrix_input,t, \n                all_atom_index, all_atom_xyz, all_atom_res_index,all_atom_valid_mask,\n                trunk_grad, N_cycle=4,res_ids=None):\n        \n        \n        sequence_features, pairwise_features=self.get_conditioning(src,distance_matrix_input,N_cycle,trunk_grad,res_ids=res_ids)\n        distogram=self.distogram_predictor(pairwise_features)\n\n\n        xyz = self.denoise(sequence_features,pairwise_features,t,\n                all_atom_index, all_atom_xyz, all_atom_res_index, all_atom_valid_mask)\n\n        return xyz, distogram\n\n\n\n    \n\n    def denoise(self,sequence_features,pairwise_features,t,\n                all_atom_index, all_atom_xyz, all_atom_res_index, all_atom_valid_mask):\n        \n        r_noisy = all_atom_xyz / torch.sqrt(self.data_std**2+t**2)[:, None, None]  # Normalize by data_std and t\n        #print(r_noisy.std())\n        L = sequence_features.shape[1]\n        all_atom_representation, all_atom_conditioning, local_attention_pair_rep, inverse_pair_distance, local_attention_pair_mask=\\\n            self.all_atom_encoder(sequence_features, pairwise_features, all_atom_index, r_noisy, all_atom_res_index, all_atom_valid_mask)\n\n        #avg_all_atom_representation = torch.stack([all_atom_representation[:,all_atom_res_index[0] == i].mean(dim=1) for i in range(L)], 1) \n\n        # c1_indices = all_atom_index.squeeze() % 30 == 5\n        # c1_xyz = r_noisy[:,c1_indices]#.reshape(config.decoder_batch_size, -1, 3)\n        #print(c1_xyz)\n        #tgt = tgt +self.xyz_norm(self.xyz_embedder(c1_xyz))\n\n        sequence_features, pairwise_features=self.get_decoder_features(sequence_features, pairwise_features, t)\n\n        #tgt=self.adaptor(sequence_features)+self.all_atom_upsample(avg_all_atom_representation)\n        #tgt=self.all_atom_upsample(avg_all_atom_representation)\n        tgt= torch.nn.functional.relu(self.all_atom_avg(all_atom_representation))\n        tgt = torch.stack([tgt[:,all_atom_res_index[0] == i].mean(dim=1) for i in range(L)], 1) \n        tgt = self.all_atom_upsample(tgt)\n\n        tgt = tgt + self.s2a(sequence_features)\n\n        for layer in self.structure_module:\n            #tgt=layer([tgt, sequence_features,pairwise_features,xyz,None])\n            tgt=checkpoint.checkpoint(self.custom(layer),\n            [tgt, sequence_features,pairwise_features,None],\n            use_reentrant=False)\n            # xyz=xyz+self.xyz_predictor(sequence_features).squeeze(0)\n            # xyzs.append(xyz)\n            #print(sequence_features.shape)\n        \n\n\n\n\n        # xyz = self.xyz_predictor(tgt).squeeze(0)\n        # return xyz\n\n        tgt = self.tgt_downsample(tgt)\n\n        all_atom_representation = all_atom_representation + tgt[:, all_atom_res_index[0]]\n\n        #always assume xyz_update has unit variance\n        r_update = self.all_atom_decoder(all_atom_representation,\n                                    all_atom_conditioning,\n                                    local_attention_pair_rep,\n                                    inverse_pair_distance,\n                                    local_attention_pair_mask)\n\n        #compute xyz_denoised\n        s_ratio = (t / self.data_std)[..., None, None].to(\n            r_update.dtype\n        )\n        x_denoised = (\n            1 / (1 + s_ratio**2) * all_atom_xyz\n            + t[..., None, None] / torch.sqrt(1 + s_ratio**2) * r_update\n        ).to(r_update.dtype)\n\n        return x_denoised\n\n\n    def extract(self, a, t, x_shape):\n        \"\"\"\n            from lucidrains' implementation\n                https://github.com/lucidrains/denoising-diffusion-pytorch/blob/beb2f2d8dd9b4f2bd5be4719f37082fe061ee450/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py#L376\n        \"\"\"\n        b, *_ = t.shape\n        out = a.gather(-1, t)\n        return out.reshape(b, *((1,) * (len(x_shape) - 1)))\n    \n    def scale_to_minus_one_to_one(self, x):\n        # according to the DDPMs paper, normalization seems to be crucial to train reverse process network\n        return x * 2 - 1\n    \n    def reverse_scale_to_zero_to_one(self, x):\n        return (x + 1) * 0.5\n    \n    def make_noisy(self, x_zeros, t): \n        # assume we get raw data, so center and scale by 35\n        x_zeros = x_zeros - torch.nanmean(x_zeros,1,keepdim=True)\n        x_zeros = x_zeros#/self.data_std\n        #rotate randomly\n        x_zeros = random_rotation_point_cloud_torch_batch(x_zeros)\n\n\n        # perturb x_0 into x_t (i.e., take x_0 samples into forward diffusion kernels)\n\n        epsilon = torch.randn_like(x_zeros).to(x_zeros.device) * t[:, None, None] # noise\n\n        noisy_sample = x_zeros + epsilon\n\n    \n        return noisy_sample.detach(), epsilon\n    \n    \n\n    \n    \n    def denoise_at_t(self, x_t, sequence_features, pairwise_features, timestep, t):\n        B, _, _ = x_t.shape\n        if t > 1:\n            z = torch.randn_like(x_t).to(sequence_features.device)\n        else:\n            z = torch.zeros_like(x_t).to(sequence_features.device)\n        \n        # at inference, we use predicted noise(epsilon) to restore perturbed data sample.\n        epsilon_pred = self.denoise(sequence_features, pairwise_features, x_t, timestep)\n        \n        alpha = self.extract(self.alphas.to(x_t.device), timestep, x_t.shape)\n        sqrt_alpha = self.extract(self.sqrt_alphas.to(x_t.device), timestep, x_t.shape)\n        sqrt_one_minus_alpha_bar = self.extract(self.sqrt_one_minus_alpha_bars.to(x_t.device), timestep, x_t.shape)\n        sqrt_beta = self.extract(self.sqrt_betas.to(x_t.device), timestep, x_t.shape)\n        \n        # denoise at time t, utilizing predicted noise\n        x_t_minus_1 = 1 / sqrt_alpha * (x_t - (1-alpha)/sqrt_one_minus_alpha_bar*epsilon_pred) + sqrt_beta*z\n        \n        return x_t_minus_1#.clamp(-1., 1)\n                \n    def sample(self, src, N, N_cycle=1):\n        device = src.device\n        x_t = torch.randn((N, src.shape[1], 3)).to(device)\n\n        # Get conditioning\n        with torch.no_grad():\n            sequence_features, pairwise_features=self.get_conditioning(src,N_cycle)\n\n        distogram = self.distogram_predictor(pairwise_features).squeeze()\n        distogram = distogram.squeeze()[:, :, 2:40].softmax(-1) * torch.arange(2, 40).float().to(device)\n        distogram = distogram.sum(-1)\n\n        for t in range(self.n_times-1, -1, -1):\n            timestep = torch.tensor([t]).repeat_interleave(N, dim=0).long().to(src.device)\n            x_t = self.denoise_at_t(x_t, sequence_features, pairwise_features, timestep, t)\n        \n        # denormalize x_0 into 0 ~ 1 ranged values.\n        #x_0 = self.reverse_scale_to_zero_to_one(x_t)\n        x_0 = x_t * self.data_std\n        return x_0, distogram\n\n    def sample_heun(self, src, N, num_steps=None, eta=0.0):\n        \"\"\"\n        Heun's method sampler with optional fewer steps (coarse stepping).\n        \n        Args:\n            src (torch.Tensor): Input sequence tensor.\n            N (int): Number of samples to generate.\n            num_steps (int, optional): Number of timesteps to sample with. If None, uses full DDPM schedule.\n        \"\"\"\n        device = src.device\n        x_t = torch.randn((N, src.shape[1], 3)).to(device)\n\n        # Get conditioning\n        # Get conditioning\n        with torch.no_grad():\n            sequence_features, pairwise_features=self.get_conditioning(src)\n\n        distogram = self.distogram_predictor(pairwise_features).squeeze()\n        distogram = distogram.squeeze()[:, :, 2:40].softmax(-1) * torch.arange(2, 40).float().to(device)\n        distogram = distogram.sum(-1)\n\n        if num_steps is None:\n            timesteps = list(range(self.n_times - 1, 0, -1))\n        else:\n            timesteps = torch.linspace(self.n_times - 1, 1, steps=num_steps, dtype=torch.long).tolist()\n        # print(timesteps)\n        # exit()\n        for i in range(len(timesteps)):\n            t = int(timesteps[i])\n            t_next = int(timesteps[i + 1]) if i + 1 < len(timesteps) else 0\n            t_next_next = int(timesteps[i + 1]) if i + 1 < len(timesteps) else 0\n\n            t_curr = torch.full((N,), t, dtype=torch.long).to(device)\n            t_next_tensor = torch.full((N,), t_next, dtype=torch.long).to(device)\n            t_next_next_tensor = torch.full((N,), t_next_next, dtype=torch.long).to(device)\n\n            alpha_bar_t = self.extract(self.sqrt_alpha_bars.to(device), t_curr, x_t.shape) ** 2\n            alpha_bar_next = self.extract(self.sqrt_alpha_bars.to(device), t_next_tensor, x_t.shape) ** 2\n            alpha_bar_next_next = self.extract(self.sqrt_alpha_bars.to(device), t_next_next_tensor, x_t.shape) ** 2\n\n            #Predict noise at x_t\n            eps1 = self.denoise(sequence_features, pairwise_features, x_t, t_curr)\n\n            # # Predict x_0\n            x_0 = (x_t - eps1 * torch.sqrt(1 - alpha_bar_t)) / torch.sqrt(alpha_bar_t)\n\n            #grad1 = eps1 * torch.sqrt(1 - alpha_bar_t)\n            # Euler step\n            x_t_euler = torch.sqrt(alpha_bar_next) * x_0 + torch.sqrt(1 - alpha_bar_next) * eps1\n            #x_t = x_t_euler\n            # # Predict noise at x_{t-1}\n\n            #print(torch.square(x_t-x_t_euler).mean())\n            #x_t_euler = batched_svd_align(x_t_euler, x_t)\n            # print(torch.square(x_t-x_t_euler).mean())\n            # exit()\n            eps2 = self.denoise(sequence_features, pairwise_features, x_t_euler, t_next_tensor)\n\n            # Predict x_0\n            #x_0 = (x_t - eps1 * torch.sqrt(1 - alpha_bar_t)) / torch.sqrt(alpha_bar_t)\n\n            x_0_next = (x_t_euler - eps2 * torch.sqrt(1 - alpha_bar_next_next)) / torch.sqrt(alpha_bar_next_next)\n            x_t_next = torch.sqrt(alpha_bar_next) * x_0_next + torch.sqrt(1 - alpha_bar_next) * eps2\n\n            x_t = 0.5 * (x_t_euler + x_t_next)\n\n\n\n        x_0 = x_t * self.data_std\n        return x_0, distogram       \n\n\n\n    def sample_euler(self, src, N, all_atom_index, all_atom_res_index, num_steps=None, eta=0.0, N_cycle=4):\n        \"\"\"\n        Heun's method sampler with optional fewer steps (coarse stepping).\n        \n        Args:\n            src (torch.Tensor): Input sequence tensor.\n            N (int): Number of samples to generate.\n            num_steps (int, optional): Number of timesteps to sample with. If None, uses full DDPM schedule.\n        \"\"\"\n        device = src.device\n        x_t = torch.randn((N, all_atom_index.shape[1], 3)).to(device)\n        # print(x_t.shape)\n        # exit()\n        # Get conditioning\n        with torch.no_grad():\n            sequence_features, pairwise_features=self.get_conditioning(src,N_cycle)\n\n        distogram = self.distogram_predictor(pairwise_features).squeeze()\n        distogram = distogram.squeeze()[:, :, 2:40].softmax(-1) * torch.arange(2, 40).float().to(device)\n        distogram = distogram.sum(-1)\n\n        if num_steps is None:\n            timesteps = list(range(self.n_times - 1, 0, -1))\n        else:\n            timesteps = torch.linspace(self.n_times - 1, 1, steps=num_steps, dtype=torch.long).tolist()\n\n        for i in range(len(timesteps)):\n            t = int(timesteps[i])\n            t_next = int(timesteps[i + 1]) if i + 1 < len(timesteps) else 0\n\n            t_curr = torch.full((N,), t, dtype=torch.long).to(device)\n            t_next_tensor = torch.full((N,), t_next, dtype=torch.long).to(device)\n\n            alpha_bar_t = self.extract(self.sqrt_alpha_bars.to(device), t_curr, x_t.shape) ** 2\n            alpha_bar_next = self.extract(self.sqrt_alpha_bars.to(device), t_next_tensor, x_t.shape) ** 2\n\n            # Predict noise at x_t\n            eps1 = self.denoise(sequence_features, pairwise_features, t_curr, all_atom_index, x_t, all_atom_res_index, None)\n\n\n            # Euler step\n            scale_factor = torch.sqrt(alpha_bar_t)/torch.sqrt(alpha_bar_next)\n            step_size = (torch.sqrt((1 - alpha_bar_next)*alpha_bar_next) / torch.sqrt(alpha_bar_t) + torch.sqrt(1 - alpha_bar_next))\n\n            #x_t_euler = x_t * scale_factor - eps1 * step_size\n\n            #x_t = x_t_euler\n\n            #if stochastic:\n            #sigma_t = eta * torch.sqrt((1 - alpha_prev) / (1 - alpha_t) * (1 - alpha_t/alpha_prev))\n            sigma_t = eta * torch.sqrt((1 - alpha_bar_next) / (1 - alpha_bar_t) * (1 - alpha_bar_t / alpha_bar_next))\n            sqrt_beta = self.extract(self.sqrt_betas.to(t_curr.device), t_curr, x_t.shape)\n            z = torch.randn_like(x_t) if i + 1 < len(timesteps) else torch.zeros_like(x_t)\n            x_t = x_t * scale_factor - eps1 * step_size +  z * sigma_t * eta\n\n\n\n\n        x_0 = x_t * self.data_std\n        return x_0, distogram        \n\n    def sample_diffusion(\n        self,\n        src,\n        distance_matrix_input,\n        N,\n        all_atom_index,\n        all_atom_res_index,\n        num_steps=200,              # Number of denoising steps\n        gamma0=0.8,\n        gamma_min=1.0,\n        lambd=1.003,\n        eta=1.5,\n        sigma_data=1.0,\n        s_max=160.0,\n        s_min=4e-4,\n        p=7.0,\n        N_cycle=4,\n    ):\n        \"\"\"\n        Implements Algorithm 18 (SampleDiffusion) with t_hat schedule computed from:\n            t_hat = σ_data * (s_max^(-1/p) + t * (s_min^(-1/p) - s_max^(-1/p)))^(-p)\n\n        Args:\n            src: Input sequence features\n            N: Batch size\n            all_atom_index: Atom indices [N, Natoms]\n            all_atom_res_index: Residue indices [N, Natoms]\n            num_steps: Number of denoising steps\n            gamma0, gamma_min, lambd, eta: Diffusion hyperparameters\n            sigma_data, s_max, s_min, p: Schedule parameters\n            N_cycle: Cycles for conditioning\n        Returns:\n            Final denoised structure [N, Natoms, 3]\n        \"\"\"\n        device = src.device\n\n        # === Step 0: compute t_hat schedule ===\n        t_values = np.linspace(0, 1, num_steps + 1)\n        s_max_1_p = s_max ** (-1.0 / p)\n        s_min_1_p = s_min ** (-1.0 / p)\n        t_hat_schedule = sigma_data * (s_max_1_p + t_values * (s_min_1_p - s_max_1_p)) ** (-p)\n        t_hat_schedule = torch.tensor(t_hat_schedule, dtype=torch.float32, device=device)  # [num_steps + 1]\n        #print(t_hat_schedule)\n        # exit()\n        # === Step 1: Initialize x_l ~ N(0, I) scaled by t_hat[0] ===\n        x_l = torch.randn((N, all_atom_index.shape[1], 3), device=device) * t_hat_schedule[0]\n\n        # === Step 2: Get conditioning ===\n        with torch.no_grad():\n            sequence_features, pairwise_features=self.get_conditioning(src,distance_matrix_input,N_cycle)\n            distogram = self.distogram_predictor(pairwise_features).squeeze()\n        distogram = distogram.squeeze()[:, :, 2:40].softmax(-1) * torch.arange(2, 40).float().to(device)\n        distogram = distogram.sum(-1)\n\n        # === Step 3: Sampling loop ===\n        for i in range(1, len(t_hat_schedule)):\n            c_tau = t_hat_schedule[i]\n            c_prev = t_hat_schedule[i - 1]\n\n            # Step 3: Centering (random augmentation)\n            x_l = x_l - x_l.mean(dim=1, keepdim=True)\n\n            # Step 4: Compute gamma\n            gamma = gamma0 if c_tau > gamma_min else 0.0\n\n            # Step 5: Compute t̂ = c_{τ−1} * (1 + γ)\n            t_hat = c_prev * (1.0 + gamma)\n\n            # Step 6: Sample noise\n            noise_scale = lambd * torch.sqrt(t_hat**2 - c_prev**2)\n            # print(t_hat)\n            # print(c_tau)\n            # print(noise_scale)\n            # exit()\n            xi = torch.randn_like(x_l) * noise_scale\n\n            # Step 7: Add noise\n            x_noisy = x_l + xi\n\n            # Step 8: Denoising\n            t_hat = torch.tensor([t_hat]).expand(N).to(device)  # [N, 1]\n            c_tau = torch.tensor([c_tau]).expand(N).to(device)  # [N, 1]\n            \n\n            x_denoised = self.denoise(\n                sequence_features,\n                pairwise_features,\n                t_hat,\n                all_atom_index,\n                x_noisy,\n                all_atom_res_index,\n                None,\n            )\n\n            t_hat = t_hat[:, None, None]  # [N, 1, 1]\n            c_tau = c_tau[:, None, None]  # [N, 1, 1]\n\n            delta = (x_noisy - x_denoised) / t_hat\n            #delta = (x_l - x_denoised) / t_hat\n            # print(torch.abs(delta2 - delta).mean())\n            # exit()\n            # Step 10: Time delta\n            dt = c_tau - t_hat\n\n            # Step 11: Euler update\n            x_l = x_noisy + eta * dt * delta\n\n        # === Step 13: Return denormalized output ===\n        return x_l, distogram","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config=load_config_from_yaml(\"/kaggle/input/rnet3d-test141/from_scratch.yaml\")\nmodel=finetuned_RibonanzaNet(load_config_from_yaml(\"/kaggle/input/ribonanzanet2/pytorch/alpha/1/pairwise.yaml\"),config,pretrained=False).cuda()\n\nimport torch\nstate_dict=torch.load(\"/kaggle/input/rnet3d-test141/from_scratch.yaml_RibonanzaNet_3D_final.pt\",map_location='cpu')\n#state_dict=torch.load(\"RibonanzaNet-3D-v2.pt\",map_location='cpu')\n\n#get rid of module. from ddp state dict\nnew_state_dict={}\n\nfor key in state_dict:\n    new_state_dict[key[7:]]=state_dict[key]\n\nmodel.load_state_dict(new_state_dict)\nmodel.eval();","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"templates=template_csv\ntemplates['target_id']=[ID.split('_')[0] for ID in templates['ID']]\ntemplates.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"preds=[]\nfor i in tqdm(range(len(test_dataset))):\n    sequence=test_data.loc[i,'sequence']\n    src=test_dataset[i]['sequence'].long()\n    all_atom_index=test_dataset[i]['all_atom_index']#.long()\n    all_atom_res_index=test_dataset[i]['all_atom_res_index'].long()\n    \n    src=src.unsqueeze(0).cuda()\n    all_atom_index=all_atom_index.unsqueeze(0).cuda()\n    all_atom_res_index=all_atom_res_index.unsqueeze(0).cuda()\n    target_id=test_data.loc[i,'target_id']\n\n    xyz=[]\n    for template_idx in range(1,5):\n        template=get_template(target_id,template_idx)\n        predicted_dm=[]\n        #for _ in range(5):\n        with torch.no_grad():\n            xyz_,distogram=model.sample_diffusion(src,template,1,\n                                                all_atom_index,all_atom_res_index,\n                                                num_steps,N_cycle=N_cycle)\n            xyz.append(xyz_)\n\n    with torch.no_grad():\n        xyz_,distogram=model.sample_diffusion(src,None,1,\n                                            all_atom_index,all_atom_res_index,\n                                            num_steps,N_cycle=N_cycle)\n        xyz.append(xyz_)\n    xyz=torch.cat(xyz,0)\n    #exit()\n    for pred_cnt in range(Npred):\n        rows=[]\n        atom_indices = []\n        for j in range(all_atom_index.shape[1]):\n            atom_index = all_atom_index[0,j][3:29].argmax().item()  # Get the index of the first non-zero atom feature\n            atom_indices.append(atom_index)\n            resname = sequence[all_atom_res_index[0,j].item()]\n            atom_name = all_atom_names_inverse[atom_index]  # Get atom name from the group\n            resid = all_atom_res_index[0,j].item() + 1  # Convert to 1-indexed\n            xyz_coord = xyz[pred_cnt,j, :].cpu().numpy()\n\n            row = [target_id, resname, atom_name, resid] + list(xyz_coord)\n            rows.append(row)\n    \n        df = pd.DataFrame(rows, columns=['target_id', 'resname', 'atom_name', 'resid', 'x', 'y', 'z'])\n        write2pdb_all_atom(df, f'{out_folder}/{target_id}_{pred_cnt}.pdb')\n    #exit()\n    atom_indices = torch.tensor(np.array(atom_indices)).cuda()\n    c1_index = atom_indices == 0\n    c1_xyz = xyz[:,c1_index,]\n    preds.append(c1_xyz.cpu().numpy())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# make sub csv\nID=[]\nresname=[]\nresid=[]\nx=[]\ny=[]\nz=[]\n\ndata=[]\n\nfor i in range(len(test_data)):\n    #print(test_data.loc[i])\n\n    \n    for j in range(len(test_data.loc[i,'sequence'])):\n        # ID.append(test_data.loc[i,'sequence_id']+f\"_{j+1}\")\n        # resname.append(test_data.loc[i,'sequence'][j])\n        # resid.append(j+1) # 1 indexed\n        row=[test_data.loc[i,'target_id']+f\"_{j+1}\",\n             test_data.loc[i,'sequence'][j],\n             j+1]\n\n        for k in range(5):\n            for kk in range(3):\n                row.append(preds[i][k][j][kk])\n        data.append(row)\n\ncolumns=['ID','resname','resid']\nfor i in range(1,6):\n    columns+=[f\"x_{i}\"]\n    columns+=[f\"y_{i}\"]\n    columns+=[f\"z_{i}\"]\n\n\nsubmission=pd.DataFrame(data,columns=columns)\n\n\nsubmission.to_csv('submission.csv',index=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}