{"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":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210},{"sourceType":"datasetVersion","sourceId":14822484,"datasetId":9479395,"databundleVersionId":15679639},{"sourceType":"datasetVersion","sourceId":11118830,"datasetId":6933267,"databundleVersionId":11511771},{"sourceType":"datasetVersion","sourceId":11922858,"datasetId":7495841,"databundleVersionId":12430654},{"sourceType":"datasetVersion","sourceId":14519720,"datasetId":9271415,"databundleVersionId":15347344},{"sourceType":"datasetVersion","sourceId":14534919,"datasetId":9283271,"databundleVersionId":15363994},{"sourceType":"datasetVersion","sourceId":14805765,"datasetId":9467172,"databundleVersionId":15661298},{"sourceType":"datasetVersion","sourceId":10855324,"datasetId":6742586,"databundleVersionId":11219268},{"sourceType":"modelInstanceVersion","sourceId":311741,"databundleVersionId":11641144,"modelInstanceId":264400,"modelId":285488}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"This notebook is a small modified version of https://www.kaggle.com/code/odat1248/rnapro-inference-with-tbm-eba .  \nFor the purpose of checking the efficiency of EBA.\n - Version 2: Only EBA (LB score: 0.307)\n - Version 3: Only Biopythons' PairwiseAligner (LB score: 0.352)\n - Version 6: EBA + Biopythons' PairwiseAligner (LB score: 0.355)\n - Version 4, 5: EBA + Biopythons' PairwiseAligner + RNAPro  (with different parameters, LB score: 0.381, 0.382)\n\nNote:\n - The multimeric targets were treated as concatenated monomer in EBA (I'm not sure about other methods).\n - If the number of hits were less than 5, other methods might be used.\n","metadata":{}},{"cell_type":"code","source":"# V2: imports recomputed ccd-cache","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-14T05:53:31.375731Z","iopub.execute_input":"2026-03-14T05:53:31.376603Z","iopub.status.idle":"2026-03-14T05:53:31.380457Z","shell.execute_reply.started":"2026-03-14T05:53:31.376566Z","shell.execute_reply":"2026-03-14T05:53:31.379535Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from Bio.Align import PairwiseAligner # to check pip installation","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T05:53:31.382563Z","iopub.execute_input":"2026-03-14T05:53:31.382932Z","iopub.status.idle":"2026-03-14T05:53:31.682089Z","shell.execute_reply.started":"2026-03-14T05:53:31.382908Z","shell.execute_reply":"2026-03-14T05:53:31.681138Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Reference\n\n- https://www.kaggle.com/code/theoviel/stanford-rna-3d-folding-pt2-rnapro-inference (by Theo Viel, Youhan Lee, chrismunley, Apache 2.0 open source license)\n- https://www.kaggle.com/code/jaejohn/rnapro-inference-with-tbm (by g john rao, Apache 2.0 open source license)","metadata":{}},{"cell_type":"code","source":"import os,sys,shutil;\nIS_SCORING_RUN = os.environ.get('KAGGLE_IS_COMPETITION_RERUN')\nprint(IS_SCORING_RUN)\nis_debug = False;\nnum_eba_predictions = 5;\nnum_tbm_predictions = 5; # Includes EBA \nsimilarity_tmscore_max = 0.8;\nsimilarity_identity_max = 999;\nskip_dot_product=True;\nskip_rnapro=True;\n\nmin_tbm_target_length = 0; # Skip TBM for short RNAs.\n\nif os.path.exists('/kaggle/input/usalign/USalign') and not os.path.exists('/kaggle/working/USalign'):\n    shutil.copy2('/kaggle/input/usalign/USalign', '/kaggle/working/USalign')\n    os.chmod('/kaggle/working/USalign', 0o755)\nusalign_bin = '/kaggle/working/USalign' if os.path.exists('/kaggle/working/USalign') else 'USalign'\n\n\"\"\"if IS_SCORING_RUN is None:\n    os.system(\"touch submission.csv\");\n    if not is_debug: # Exit to save GPU Time\n        quit();\"\"\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T05:53:31.683287Z","iopub.execute_input":"2026-03-14T05:53:31.684107Z","iopub.status.idle":"2026-03-14T05:53:31.724344Z","shell.execute_reply.started":"2026-03-14T05:53:31.684078Z","shell.execute_reply":"2026-03-14T05:53:31.723665Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Prepare EBA","metadata":{}},{"cell_type":"code","source":"import re,os,sys,gzip,tempfile,shutil,stat,random,copy,subprocess,json,numpy as np;\n# Skip sequences with large differences in length\nmax_length_diff = 0.5;\n\ntraining_rep_dataset = \"/kaggle/input/datasets/odat1248/stanfordrna2026-training-ribonanza-rep\";\nribonanzanet_dataset = \"/kaggle/input/datasets/odat1248/ribonanzanet-custom-set\";\nmatrix_align_dataset = \"/kaggle/input/datasets/odat1248/matrix-align-20260207-avx2\";\n\nribonanza_net_script = ribonanzanet_dataset+\"/ribonanzanet_process.py\"\nribonanza_net_config = ribonanzanet_dataset+\"/RibonanzaNet_custom/RibonanzaNet/configs/pairwise.yaml\";\nribonanza_net_model = ribonanzanet_dataset+\"/archive/RibonanzaNet.pt\";\n\ntraining_full_rep_dir = training_rep_dataset+\"/StanfordRNA2026_training_ribonanza_rep/train_full_ribonanza_rep\";\ntraining_full_json_dir = training_rep_dataset+\"/StanfordRNA2026_training_ribonanza_rep/jsons\";\ntrain_full_fasta = training_rep_dataset+\"/StanfordRNA2026_training_ribonanza_rep/train_full.fasta\";\ntrain_mask_fasta = training_rep_dataset+\"/StanfordRNA2026_training_ribonanza_rep/train_modeled_mask.fasta\";\nmatrix_align_binary = matrix_align_dataset+\"/matrix_align\";\n\nshutil.copy2(matrix_align_binary,\"./\");\nmatrix_align_binary = \"./matrix_align\";\nos.chmod(matrix_align_binary,stat.S_IXUSR+stat.S_IRUSR+stat.S_IXGRP+stat.S_IRGRP+stat.S_IXOTH+stat.S_IROTH);\n\nprevpath = \"\";\nif \"PYTHONPATH\" in os.environ:\n    prevpath = str(os.environ[\"PYTHONPATH\"]);\nrpath = ribonanzanet_dataset+\"/RibonanzaNet_custom/RibonanzaNet\";\nrpath = re.sub(r\"/+\",\"/\",rpath);\nif rpath not in prevpath:\n    if len(prevpath) == 0:\n        os.environ[\"PYTHONPATH\"] = rpath;\n    else:\n        os.environ[\"PYTHONPATH\"] = rpath+\":\"+prevpath;\n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T05:53:31.725402Z","iopub.execute_input":"2026-03-14T05:53:31.725703Z","iopub.status.idle":"2026-03-14T05:53:31.796081Z","shell.execute_reply.started":"2026-03-14T05:53:31.725669Z","shell.execute_reply":"2026-03-14T05:53:31.79537Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def loadFasta(filename):\n    if filename.endswith(\".gz\"):\n        import gzip;\n        fin = gzip.open(filename,\"rt\");\n    else:\n        fin = open(filename,\"rt\");\n    ret = [];\n    cdict = dict();\n    cdict[\"seq\"] = \"\";\n    ret.append(cdict);\n    \n    for ll in fin:\n        mat = re.search(r\"[\\s]*>\",ll);\n        if mat is not None:\n            cdict = dict();\n            ret.append(cdict);\n            nmat = re.search(r\"[\\s]*>[\\s]*([^\\s]+)\",ll);\n            if nmat is not None:\n                cdict[\"name\"] = nmat.group(1);\n                cdict[\"desc\"] = \"\"\n            dmat = re.search(r\"[\\s]*>[\\s]*([^\\s]+)[\\s]+([^\\s][^\\r\\n]*)\",ll);\n            if dmat is not None:\n                cdict[\"desc\"] = dmat.group(2);\n            cdict[\"seq\"] = \"\";\n        else:\n            cdict[\"seq\"] += re.sub(r\"[\\s]\",\"\",ll);\n            \n    if len(ret[0][\"seq\"]) == 0:\n        ret.pop(0);\n    fin.close();\n    return ret;\n\ndef line_to_dict(line):\n    line = re.sub(r\"[\\r\\n]\",\"\",line);\n    line = re.sub(r\"[\\s]*:[\\s*]*\",\":\",line);\n    pt = re.split(r\"[\\t]+\",line);\n    ret = {};\n    ret[\"nokey\"] = [];\n    for p in pt:\n        if len(p) == 0:\n            continue;\n        mat = re.search(r\"^([^:]+):(.+)$\",p);\n        if mat is None:\n            ret[\"nokey\"].append(p);\n        else:\n            k = mat.group(1);\n            v = mat.group(2);\n            if k in ret:\n                sys.stderr.write(\"Duplicated key: \"+k+\"\\n\");\n                sys.stderr.write(v+\" was ignored.\\n\");\n                continue;\n            else:\n                ret[k] = v;\n    return ret;\n\ndef a3m_to_mfa(query_,target_):\n    if re.search(r\"[a-z]\",query_):\n        sys.stderr.write(\"Lowercase letter is not expected... but process anyway...\\n\");\n    query = re.sub(r\"[\\s]\",\"\",query_);\n    target = re.sub(r\"[\\s]\",\"\",target_);\n    q = re.sub(r\"[\\s]\",\"\",query);\n    t = re.sub(r\"[a-z\\s]\",\"\",target);\n    assert len(q) == len(t);\n    \n    query = list(query);\n    target = list(target);\n    \n    apos = -1;\n    bpos = -1;\n    ret_q = [];\n    ret_t = [];\n    \n    for tt in target:\n        if tt == \"-\":\n            apos += 1;\n            ret_q.append(query[apos]);\n            ret_t.append(tt);\n            continue;\n        if tt.islower():\n            ret_q.append(\"-\");\n            ret_t.append(tt.upper());\n            continue;\n        apos += 1;\n        ret_q.append(query[apos]);\n        ret_t.append(tt);\n    return \"\".join(ret_q),\"\".join(ret_t);\n    \ndef get_first_line(file):\n    if file.endswith(\".gz\"):\n        fin = gzip.open(file,\"rt\");\n    else:\n        fin = open(file,\"rt\");\n    l = fin.readline();\n    fin.close();\n    return l;","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T05:53:31.798313Z","iopub.execute_input":"2026-03-14T05:53:31.798819Z","iopub.status.idle":"2026-03-14T05:53:31.814063Z","shell.execute_reply.started":"2026-03-14T05:53:31.798794Z","shell.execute_reply":"2026-03-14T05:53:31.813188Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_atom_line(serial_number,atom_name,residue_name,chain_id,residue_pos,xx,yy,zz,element):\n    \n    xx = \"{:>.3f}\".format(xx);\n    yy = \"{:>.3f}\".format(yy);\n    zz = \"{:>.3f}\".format(zz);\n    \n    if len(xx) > 8:\n        raise Exception(\"string overflow\".format(self.x));\n    if len(yy) > 8:\n        raise Exception(\"string overflow\".format(self.y));\n    if len(zz) > 8:\n        raise Exception(\"string overflow\".format(self.z));\n        \n    ret = \"{head:<6}{serial_number:>5} {atom_name:<4}{alt_loc:1}{residue_name:>3} {chain_id:1}{residue_pos:>4}{insertion_code:1}   {xx:>8}{yy:>8}{zz:>8}{occupancy:>6}{bfactor:>6}          {element:>2} {charge:2}\".format(        \n    head = \"ATOM  \",\n    serial_number = serial_number,\n    atom_name = atom_name,\n    alt_loc = \" \",\n    residue_name = residue_name,\n    chain_id = chain_id,\n    residue_pos = residue_pos,\n    insertion_code = \" \",\n    xx = xx,\n    yy = yy,\n    zz = zz,\n    occupancy = 1.0,\n    bfactor = 0.0,\n    element = element,\n    charge = 0.0);\n    return ret;\ndef to_pdb(a,outfile,aseq = None):\n    with open(outfile, \"wt\") as aout:\n        for i in range(a.shape[0]):\n            if aseq is None:\n                resname = \"A\";\n            else:\n                resname = aseq[i];\n            aout.write(\n                create_atom_line(serial_number=i,atom_name=\" C1'\",residue_name=resname,chain_id=\"A\"\n                                 ,residue_pos=i,xx=a[i,0],yy=a[i,1],zz=a[i,1],element=\"C\")\n            );\n            aout.write(\"\\n\");\ndef get_tmscore(a,b,aseq=None,bseq=None,tmalign=False):\n    assert len(a.shape) == 2, a.shape[-1] == 3;\n    assert len(b.shape) == 2, abshape[-1] == 3;\n    with tempfile.TemporaryDirectory() as tmpdirname:\n        afile = os.path.join(tmpdirname,\"a.pdb\");\n        bfile = os.path.join(tmpdirname,\"b.pdb\");\n        to_pdb(a,afile,aseq);\n        to_pdb(b,bfile,bseq);\n        proc = subprocess.run([usalign_bin,afile,bfile,\"-TMscore\",\"0\" if tmalign else \"1\",\"-atom\",\" C1'\",\"-outfmt\",\"2\"], capture_output=True, text=True);\n        lin = re.split(r\"[\\r\\n]\",proc.stdout);\n        for i in range(len(lin)):\n            pt = re.split(r\"[\\s]+\",lin[i]);\n            if len(pt) > 2:\n                if \"TM1\" == pt[2]:\n                    pt = re.split(r\"[\\s]+\",lin[i+1]);\n                    return float(pt[2]);\n        sys.stderr.write(\"USalign Failed\");\n        sys.stderr.write(proc.stderr);\n        sys.stderr.write(proc.stdout);\n    return -1.0;","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T05:53:31.815102Z","iopub.execute_input":"2026-03-14T05:53:31.815356Z","iopub.status.idle":"2026-03-14T05:53:31.834616Z","shell.execute_reply.started":"2026-03-14T05:53:31.815335Z","shell.execute_reply":"2026-03-14T05:53:31.833713Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"rep_training_seqs = {};\nfor s in os.listdir(training_full_json_dir):\n    \n    jsonpath = os.path.join(training_full_json_dir,s);\n\n    name = re.sub(r\"\\.json(\\.gz)?$\",\"\",s);\n    rep_training_seqs[name] = {};\n    \n    assert os.path.exists(jsonpath);\n    rep_training_seqs[name][\"json\"] = jsonpath;\n    \n    with open(jsonpath) as fin:\n        j = json.load(fin);\n        \n    for k in [\"all_bases\",\"modeled_mask\"]:\n        rep_training_seqs[name][k] = j[k];\n    \n    matpath = training_full_rep_dir+\"/\"+name+\".mat\";\n    assert os.path.exists(matpath), matpath+\" does not exists.\";\n    rep_training_seqs[name][\"mat\"] = matpath;\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T05:53:31.83585Z","iopub.execute_input":"2026-03-14T05:53:31.83618Z","iopub.status.idle":"2026-03-14T05:54:36.113431Z","shell.execute_reply.started":"2026-03-14T05:53:31.836139Z","shell.execute_reply":"2026-03-14T05:54:36.112473Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_mat_file(fname):\n    if fname.endswith(\".gz\"):\n        fin = gzip.open(fname,\"rt\");\n    else:\n        fin = open(fname,\"rt\");\n    lines = fin.readlines();\n    fin.close();\n    l = lines[0];\n    mat = re.search(r\">([^\\s]+)\",l);\n    if mat:\n        mname = mat.group(1);\n    else:\n        raise Exception(\"???\");\n\n    val = [];\n    for l in lines[1:]:\n        if l.startswith(\"//\"):\n            continue;\n        if l.startswith(\"#\"):\n            continue;\n        l = re.sub(r\"[\\r\\n]\",\"\",l);\n        if len(l) == 0:\n            continue;\n        pt = [float(x) for x in re.split(r\"[\\s]+\",l)[1:]];\n        val.append(pt);\n        \n    val = np.array(val,dtype=np.float32);\n    \n    matmean = np.mean(val,axis=0);\n    \n    z = np.sum(matmean*matmean);\n    if z > 0.0:\n        matmean /= np.sqrt(z);\n        \n    return mname, val, matmean;\n\ndef run_eba(query_seq,use_dot_product=False,num_prefilter_seqs=5):\n    global matname_to_seqname,emb_mat_files,max_length_diff,is_debug;\n    print(\"Run EBA: \"+query_seq);\n    qlen = len(query_seq);\n    with tempfile.TemporaryDirectory() as tmpdirname:\n        scoreout = os.path.join(tmpdirname,\"tmp.\"+str(random.random())+\".score\");\n        tmp_template_list = os.path.join(tmpdirname,\"tmp.\"+str(random.random())+\".list\");\n        if use_dot_product:\n            matrix_align_options = [\n                matrix_align_binary,\n                \"--in_list\",\n                tmp_template_list,\n                \"--a3m_pairwise\",\n                \"true\",\n                \"--alignment_type\",\n                \"local\",\n                \"--score_type\",\n                \"dot_product\",\n                \"--gap_open_penalty\",\n                \"-604.420390841876\",\n                \"--gap_extension_penalty\",\n                \"-519.801536124013\",\n                \"--gap_penalty_auto_adjust\",\n                \"false\",\n                \"--num_threads\",\n                \"2\",\n                \"--score_out\",\n                scoreout\n            ];\n        else:\n            matrix_align_options = [\n                matrix_align_binary,\n                \"--in_list\",\n                tmp_template_list,\n                \"--a3m_pairwise\",\n                \"true\",\n                \"--alignment_type\",\n                \"local\",\n                \"--score_type\",\n                \"pearson_correl\",\n                \"--gap_open_penalty\",\n                \"-1.7\",\n                \"--gap_extension_penalty\",\n                \"-0.187\",\n                \"--gap_penalty_auto_adjust\",\n                \"false\",\n                \"--num_threads\",\n                \"2\",\n                \"--score_out\",\n                scoreout\n            ];\n        ribonanza_net_options=[\n            \"python\",\n            ribonanza_net_script,\n            \"--config_path\",\n            ribonanza_net_config,\n            \"--use_chunk\",\n            \"True\",\n            \"--split_len\",\n            \"512\",\n            \"--model_pt_path\",\n            ribonanza_net_model,\n            \"--max_len\",\n            \"10000\"\n        ];\n        \n        tmpname = os.path.join(tmpdirname,\"tmp.\"+str(random.random())+\".fas\");\n        with open(tmpname,\"wt\") as fout:\n            fout.write(\">dummy_query\\n\");\n            fout.write(query_seq+\"\\n\");\n            \n        comm = copy.deepcopy(ribonanza_net_options);\n        comm.append(\"--in_fasta\");\n        comm.append(tmpname);\n        comm.append(\"--out_dir\");\n        comm.append(tmpname+\"_res\");\n        tmpout = tmpname+\"_res/seq_0.mat.gz\";\n\n        proc = subprocess.run(comm);\n\n        _, tt, pp = load_mat_file(tmpout);\n        normalized_emb = per_sequence_norm(pp);\n        \n        topx = min(len(emb_mat_files),num_prefilter_seqs);\n            \n        # Filter embeddings with pooled value similarity\n        file_score = [];\n        for fname in list(sorted(emb_mat_files.keys())):\n            if min(qlen,emb_mat_files[fname][\"length\"]) < max(qlen,emb_mat_files[fname][\"length\"])*max_length_diff:\n                continue;\n            val = emb_mat_files[fname][\"per_seq_emb\"];\n            score = np.sum(val*normalized_emb); # dot product or cosine similarity\n            file_score.append([fname,score]);\n            \n        if len(file_score) == 0:\n            sys.stderr.write(\"Query \"+query_seq+\" does not have template with similar sequence length.\\n\");\n            return [];\n            \n        file_score_ = list(sorted(file_score,key=lambda x:x[1],reverse=True));\n        \n        # skip similar sequences\n        file_score = [];\n        retained_sequences = [];\n        for fname,score in file_score_:\n            remflag = False;\n            aseq = emb_mat_files[fname][\"sequence\"];\n            \n            for bseq in list(retained_sequences):\n                alignres = ALIGNER.align(aseq,bseq);\n                \n                matchcount = 0;\n                for a,b in zip(alignres[0].indices[0],alignres[0].indices[1]):\n                    if a != -1 and b != -1:\n                        if aseq[a] == bseq[b]:\n                            matchcount += 1;\n                            \n                if is_debug:\n                    print(\"Alignment Result:\",matchcount/len(aseq),aseq,bseq);\n                if matchcount/len(aseq) > similarity_identity_max:\n                    remflag = True;\n                    break;\n            if remflag:\n                continue;\n            \n            retained_sequences.append(aseq);\n            file_score.append((fname,score));\n            if len(file_score) >= topx:\n                break;\n        print(\"Top 5 embeddings with pooled value similarity:\",file_score[:5])\n        \n        with open(tmp_template_list,\"wt\") as fout:\n            for f in list(file_score[:topx]):\n                fout.write(f[0]+\"\\n\");\n                \n        alignres=tmpname+\".align.res\";\n        \n        comm = copy.deepcopy(matrix_align_options);\n        comm.append(\"--in\");\n        comm.append(tmpout);\n        comm.append(\"--out\");\n        comm.append(alignres);\n        \n        proc = subprocess.run(comm,stdout=subprocess.PIPE,encoding=\"utf-8\");\n    \n        name_to_score = {};\n    \n        with open(scoreout,\"rt\") as fin:\n            for l in fin:\n                d = line_to_dict(l);\n                name_to_score[d[\"sname\"]] = float(d[\"score\"]);\n                \n        hitseq = loadFasta(alignres);\n        res = [];\n        for hh in list(hitseq):\n            if hh[\"name\"] == \"dummy_query\":\n                continue;\n            hh[\"score\"] = name_to_score[hh[\"name\"]];\n            if hh[\"name\"] not in matname_to_seqname:\n                if not is_debug:\n                    sys.stderr.write(hh[\"name\"] + \" in mat file is not found in sequence list.\\n\");\n                continue;\n                    \n            hh[\"name\"] = matname_to_seqname[hh[\"name\"]];\n            res.append(hh);\n        return res;","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T05:54:36.114603Z","iopub.execute_input":"2026-03-14T05:54:36.115Z","iopub.status.idle":"2026-03-14T05:54:36.137977Z","shell.execute_reply.started":"2026-03-14T05:54:36.114973Z","shell.execute_reply":"2026-03-14T05:54:36.136951Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"matname_to_seqname = {};\n\nmatfilename_to_pooled = {};\nmatfilename_to_sequence = {};\nnormvalues = [];\nfor sname in list(sorted(rep_training_seqs.keys())):\n    r = rep_training_seqs[sname];\n    assert \"mat\" in r;\n    fname = r[\"mat\"];\n    mname, tt, pp = load_mat_file(fname);\n    matfilename_to_sequence[fname] = r[\"all_bases\"];\n    matfilename_to_pooled[fname] = pp;\n    normvalues.append(pp);\n    matname_to_seqname[mname] = sname;\n    if is_debug and len(matname_to_seqname) > 200:\n        break;\n\nnorms_mean = np.mean(normvalues,axis=0);\nnorms_std = np.std(normvalues,axis=0);\n\ndef per_sequence_norm(v_):\n    global norms_mean,norms_std;\n    v = v_ - norms_mean\n    v /= norms_std+1.0e-8\n    vs = (v*v).sum();\n    if vs > 0.0:\n        v /= np.sqrt(vs);\n    return v;\nemb_mat_files = {};\nfor k in list(sorted(matfilename_to_pooled.keys())):\n    emb_mat_files[k] = {};\n    emb_mat_files[k][\"per_seq_emb\"] = per_sequence_norm(matfilename_to_pooled[k]);\n    emb_mat_files[k][\"length\"] = len(matfilename_to_sequence[k]); # for backward compatibility\n    emb_mat_files[k][\"sequence\"] = matfilename_to_sequence[k];\ndel matfilename_to_pooled;","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T05:54:36.139137Z","iopub.execute_input":"2026-03-14T05:54:36.139493Z","iopub.status.idle":"2026-03-14T06:05:55.866123Z","shell.execute_reply.started":"2026-03-14T05:54:36.139455Z","shell.execute_reply":"2026-03-14T06:05:55.865409Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def sort_matrix_align_hit(aligned_hits):\n    global rep_training_seqs;\n    ret = [];\n    min_score = aligned_hits[0][\"score\"]\n    for hit in aligned_hits:\n        if hit[\"name\"] == \"dummy_query\":\n            continue;\n        min_score = min(min_score,hit[\"score\"])\n        if hit[\"name\"] not in rep_training_seqs:\n            sys.stderr.write(hit[\"name\"]+\" is not found in template list.\\n\");\n            continue;\n        r = rep_training_seqs[hit[\"name\"]];\n        \n        for bk in [\"all_bases\",\"modeled_mask\",\"json\"]:\n            hit[bk] = r[bk];\n            \n        mask_seq = r[\"modeled_mask\"];\n        modeled_index = {};\n        for i,s in enumerate(list(mask_seq)):\n            if s != \"-\":\n                modeled_index[i] = s;\n        apos = -1;\n        uppercount = 0;\n        modeledcount = 0;\n        for i,s in enumerate(list(hit[\"seq\"])):\n            if s == \"-\":\n                continue;\n            apos += 1;\n            if s.isupper():\n                uppercount += 1;\n                if apos in modeled_index:\n                    modeledcount += 1;\n                    if modeled_index[apos] != \"N\" and modeled_index[apos] != \"X\":\n                        assert modeled_index[apos] == s, \"???? Position \"+str(apos)+\"\\n\"+mask_seq+\" vs \"+hit[\"seq\"];\n        if uppercount != 0:\n            hit[\"modeled_ratio\"] = modeledcount/uppercount;\n        else:\n            hit[\"modeled_ratio\"] = 0.0;\n    for hit in aligned_hits:\n        if \"modeled_ratio\" in hit:\n            hit[\"modeled_score\"] = hit[\"modeled_ratio\"]*(hit[\"score\"]-min_score);\n            ret.append((hit[\"modeled_score\"],hit));\n    return [y[1] for y in sorted(ret,key=lambda x:x[0],reverse=True)];","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T06:05:55.867313Z","iopub.execute_input":"2026-03-14T06:05:55.867622Z","iopub.status.idle":"2026-03-14T06:05:55.877038Z","shell.execute_reply.started":"2026-03-14T06:05:55.86759Z","shell.execute_reply":"2026-03-14T06:05:55.876387Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp -r /kaggle/input/rnapro-src/RNAPro .\n!cp /kaggle/input/rnapro-src/rnapro-private-best-500m.ckpt .","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T06:05:55.877899Z","iopub.execute_input":"2026-03-14T06:05:55.878135Z","iopub.status.idle":"2026-03-14T06:06:38.638863Z","shell.execute_reply.started":"2026-03-14T06:05:55.878106Z","shell.execute_reply":"2026-03-14T06:06:38.637872Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cd RNAPro","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T06:06:38.640255Z","iopub.execute_input":"2026-03-14T06:06:38.640688Z","iopub.status.idle":"2026-03-14T06:06:38.64738Z","shell.execute_reply.started":"2026-03-14T06:06:38.640638Z","shell.execute_reply":"2026-03-14T06:06:38.64645Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pwd","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T06:06:38.650699Z","iopub.execute_input":"2026-03-14T06:06:38.651231Z","iopub.status.idle":"2026-03-14T06:06:38.663197Z","shell.execute_reply.started":"2026-03-14T06:06:38.65121Z","shell.execute_reply":"2026-03-14T06:06:38.662167Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pip install -e . --no-deps","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T06:06:38.664472Z","iopub.execute_input":"2026-03-14T06:06:38.664863Z","iopub.status.idle":"2026-03-14T06:06:46.464127Z","shell.execute_reply.started":"2026-03-14T06:06:38.664834Z","shell.execute_reply":"2026-03-14T06:06:46.462852Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cd ..","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T06:06:46.465701Z","iopub.execute_input":"2026-03-14T06:06:46.466021Z","iopub.status.idle":"2026-03-14T06:06:46.471721Z","shell.execute_reply.started":"2026-03-14T06:06:46.465989Z","shell.execute_reply":"2026-03-14T06:06:46.471066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pwd","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T06:06:46.472846Z","iopub.execute_input":"2026-03-14T06:06:46.473202Z","iopub.status.idle":"2026-03-14T06:06:48.301489Z","shell.execute_reply.started":"2026-03-14T06:06:46.473179Z","shell.execute_reply":"2026-03-14T06:06:48.300596Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"TBM\n---\n","metadata":{}},{"cell_type":"code","source":"# Cell 1: Imports and Setup\nimport os\nimport pandas as pd\nimport numpy as np\nfrom scipy.spatial.transform import Rotation as R\nfrom scipy.spatial import distance_matrix\nimport random\nimport time\nimport warnings\n\nfrom tqdm.auto import tqdm\nwarnings.filterwarnings('ignore')\nfrom Bio.Align import PairwiseAligner\nfrom Bio.Seq import Seq\n\n# Enable tqdm for pandas operations\ntqdm.pandas()\n    \n# Initialize global aligner with RNA-appropriate parameters\nALIGNER = PairwiseAligner()\nALIGNER.mode = 'global'\nALIGNER.match_score = 2.1\nALIGNER.mismatch_score = -1\nALIGNER.open_gap_score = -10\nALIGNER.extend_gap_score = -0.5\n\nseed = 21\nnp.random.seed(seed)\nrandom.seed(seed)\n\n# Cell 2: Load Data\nBASE_PATH = '/kaggle/input/stanford-rna-3d-folding-2'\n\n# Load with progress indication\ntest_seqs = pd.read_csv(f'{BASE_PATH}/test_sequences.csv')\nif is_debug:\n    test_seqs = test_seqs[:5];\ntrain_seqs = pd.read_csv(f'{BASE_PATH}/train_sequences.csv')\nvalidation_seqs = pd.read_csv(f'{BASE_PATH}/validation_sequences.csv')\ntrain_labels = pd.read_csv(f'{BASE_PATH}/train_labels.csv', low_memory=False)\nvalidation_labels = pd.read_csv(f'{BASE_PATH}/validation_labels.csv')\nsample_submission = pd.read_csv(f'{BASE_PATH}/sample_submission.csv')\n\nprint(f\"✓ Loaded {len(train_seqs)} training sequences\")\nprint(f\"✓ Loaded {len(validation_seqs)} validation sequences\") \nprint(f\"✓ Loaded {len(test_seqs)} test sequences\")\n\n# Cell 3: Process Training Labels to Coordinate Dictionary\ndef process_labels(labels_df, use_first_model_only=True):\n    \"\"\"\n    Process labels dataframe to create a dictionary mapping target_id to coordinates.\n    Vectorized implementation for improved performance.\n    \n    Args:\n        labels_df: DataFrame with ID, resid, x_1, y_1, z_1, etc.\n        use_first_model_only: If True, only extract first model coordinates\n        \n    Returns:\n        Dictionary mapping target_id to numpy array of coordinates\n    \"\"\"\n    print(\"Extracting target IDs...\")\n    # Vectorized target_id extraction\n    labels_df = labels_df.copy()\n    labels_df['target_id'] = labels_df['ID'].str.rsplit('_', n=1).str[0]\n    \n    # Sort once for all groups\n    labels_df = labels_df.sort_values(['target_id', 'resid'])\n    \n    # Vectorized coordinate extraction\n    print(\"Extracting coordinates...\")\n    coord_cols = ['x_1', 'y_1', 'z_1']\n    \n    # Replace placeholder values with NaN in one operation\n    coords_array = labels_df[coord_cols].values.copy()\n    coords_array[coords_array < -1e6] = np.nan\n    labels_df[coord_cols] = coords_array\n    \n    # Group and convert to dictionary with progress bar\n    print(\"Grouping by target_id...\")\n    coords_dict = {}\n    \n    grouped = labels_df.groupby('target_id', sort=False)\n    for target_id, group in tqdm(grouped, desc=\"Processing structures\", total=len(grouped)):\n        # Directly extract coordinates as numpy array (already sorted)\n        coords_dict[target_id] = group[coord_cols].values\n    \n    return coords_dict\n\n\ntrain_coords_dict = process_labels(train_labels)\nvalid_coords_dict = process_labels(validation_labels)\n\n# Cell 4: Sequence Alignment and Template Finding\ndef get_alignment_score(query_seq, template_seq):\n    alignments = ALIGNER.align(query_seq, template_seq)\n    best_alignment = next(iter(alignments), None)\n    \n    if best_alignment is None:\n        return 0.0\n    \n    # Normalize score by theoretical max (perfect match of shorter sequence)\n    max_possible = 2.1 * min(len(query_seq), len(template_seq))\n    normalized_score = best_alignment.score / max_possible\n    \n    return min(normalized_score, 1.0)  # Cap at 1.0\n\n\ndef get_aligned_sequences(query_seq, template_seq):\n    alignments = ALIGNER.align(query_seq, template_seq)\n    best_alignment = next(iter(alignments), None)\n    \n    if best_alignment is None:\n        return None, None\n    \n    # Use the robust method that builds from alignment coordinates\n    return build_aligned_sequences(query_seq, template_seq, best_alignment)\n\n\ndef build_aligned_sequences(query_seq, template_seq, alignment):\n    \"\"\"\n    Build aligned sequences with gaps from alignment coordinates.\n    Uses alignment.aligned property which returns tuples of (start, end) ranges.\n    \"\"\"\n    # Get aligned blocks: alignment.aligned returns ((query_ranges), (template_ranges))\n    query_ranges, template_ranges = alignment.aligned\n    \n    # If no aligned blocks, return None\n    if len(query_ranges) == 0:\n        return None, None\n    \n    aligned_query = []\n    aligned_template = []\n    \n    query_pos = 0\n    template_pos = 0\n    \n    for (q_start, q_end), (t_start, t_end) in zip(query_ranges, template_ranges):\n        # Add gaps for unaligned query residues (query has residues, template doesn't)\n        while query_pos < q_start:\n            aligned_query.append(query_seq[query_pos])\n            aligned_template.append('-')\n            query_pos += 1\n        \n        # Add gaps for unaligned template residues (template has residues, query doesn't)\n        while template_pos < t_start:\n            aligned_query.append('-')\n            aligned_template.append(template_seq[template_pos])\n            template_pos += 1\n        \n        # Add aligned region (both have residues)\n        block_len = q_end - q_start  # Should equal t_end - t_start\n        for i in range(block_len):\n            aligned_query.append(query_seq[q_start + i])\n            aligned_template.append(template_seq[t_start + i])\n        \n        query_pos = q_end\n        template_pos = t_end\n    \n    # Add any remaining unaligned query residues at the end\n    while query_pos < len(query_seq):\n        aligned_query.append(query_seq[query_pos])\n        aligned_template.append('-')\n        query_pos += 1\n    \n    # Add any remaining unaligned template residues at the end\n    while template_pos < len(template_seq):\n        aligned_query.append('-')\n        aligned_template.append(template_seq[template_pos])\n        template_pos += 1\n    \n    return ''.join(aligned_query), ''.join(aligned_template)\n\n\ndef find_similar_sequences(query_seq, train_seqs_df, train_coords_dict, \n                          temporal_cutoff=None, top_n=5):\n    \"\"\"\n    Find sequences in the training data similar to the query sequence.\n    Uses modern PairwiseAligner for sequence comparison.\n    \n    Args:\n        query_seq: The RNA sequence to find templates for\n        train_seqs_df: DataFrame containing training sequences\n        train_coords_dict: Dictionary mapping target_ids to their 3D coordinates\n        temporal_cutoff: Only consider training sequences published before this date\n        top_n: Number of top templates to return\n        \n    Returns:\n        List of (target_id, sequence, similarity_score, coordinates) tuples\n    \"\"\"\n    similar_seqs = []\n    \n    # Filter training sequences by temporal cutoff if provided\n    if temporal_cutoff:\n        filtered_train_seqs = train_seqs_df[train_seqs_df['temporal_cutoff'] < temporal_cutoff]\n    else:\n        filtered_train_seqs = train_seqs_df\n    \n    for _, row in filtered_train_seqs.iterrows():\n        target_id = row['target_id']\n        train_seq = row['sequence']\n        \n        # Skip if coordinates not available\n        if target_id not in train_coords_dict:\n            continue\n        \n        # Skip if sequence length difference is too large (>50%)\n        len_diff = abs(len(train_seq) - len(query_seq)) / max(len(train_seq), len(query_seq))\n        if len_diff > 0.5:\n            continue\n        \n        # Calculate similarity score using new aligner\n        similarity_score = get_alignment_score(query_seq, train_seq)\n        \n        similar_seqs.append((target_id, train_seq, similarity_score, train_coords_dict[target_id]))\n    \n    # Sort by similarity score (higher is better)\n    similar_seqs.sort(key=lambda x: x[2], reverse=True)\n    return similar_seqs[:top_n]\n\n# Cell 5: Template Adaptation Functions\ndef adapt_template_to_query(query_seq, template_seq, template_coords,precomputed_alignment=None):\n    # Get aligned sequences\n    if precomputed_alignment is not None:\n        aligned_query, aligned_template = precomputed_alignment;\n    else:\n        aligned_query, aligned_template = get_aligned_sequences(query_seq, template_seq)\n    \n    if aligned_query is None or aligned_template is None:\n        return adapt_template_simple(query_seq, template_seq, template_coords)\n    \n    # Initialize coordinates for query sequence\n    query_coords = np.zeros((len(query_seq), 3))\n    query_coords.fill(np.nan)\n    \n    # Map template coordinates to query based on alignment\n    query_idx = 0\n    template_idx = 0\n    \n    for i in range(len(aligned_query)):\n        query_char = aligned_query[i]\n        template_char = aligned_template[i]\n        \n        if query_char != '-' and template_char != '-':\n            # Both aligned - copy template coordinate to query\n            if query_idx < len(query_seq) and template_idx < len(template_coords):\n                # Handle NaN coordinates in template\n                if not np.any(np.isnan(template_coords[template_idx])):\n                    query_coords[query_idx] = template_coords[template_idx]\n            template_idx += 1\n            query_idx += 1\n        elif query_char != '-' and template_char == '-':\n            # Gap in template - query residue has no template coord\n            query_idx += 1\n        elif query_char == '-' and template_char != '-':\n            # Gap in query - skip template residue\n            template_idx += 1\n    \n    # Fill in gaps by interpolation\n    query_coords = fill_coordinate_gaps(query_coords)\n    \n    return query_coords\n\n\ndef adapt_template_simple(query_seq, template_seq, template_coords):\n    \"\"\"Simple template adaptation without Biopython alignment.\"\"\"\n    query_coords = np.zeros((len(query_seq), 3))\n    \n    # Simple mapping based on position\n    scale = len(template_coords) / len(query_seq)\n    for i in range(len(query_seq)):\n        template_idx = int(i * scale)\n        template_idx = min(template_idx, len(template_coords) - 1)\n        if not np.any(np.isnan(template_coords[template_idx])):\n            query_coords[i] = template_coords[template_idx]\n        else:\n            query_coords[i] = [np.nan, np.nan, np.nan]\n    \n    # Fill gaps\n    query_coords = fill_coordinate_gaps(query_coords)\n    return query_coords\n\n\ndef fill_coordinate_gaps(coords):\n    \"\"\"Fill NaN gaps in coordinates by interpolation.\"\"\"\n    n = len(coords)\n    typical_step = 4.0\n    \n    # First pass: interpolate between valid points\n    for i in range(n):\n        if np.isnan(coords[i, 0]):\n            prev_valid = next((j for j in range(i-1, -1, -1) if not np.isnan(coords[j, 0])), -1)\n            next_valid = next((j for j in range(i+1, n) if not np.isnan(coords[j, 0])), -1)\n            \n            if prev_valid >= 0 and next_valid >= 0:\n                weight = (i - prev_valid) / (next_valid - prev_valid)\n                coords[i] = (1 - weight) * coords[prev_valid] + weight * coords[next_valid]\n    \n    # Second pass: handle remaining NaNs at edges\n    for i in range(n):\n        if np.isnan(coords[i, 0]):\n            if i == 0:\n                first_valid = next((j for j in range(1, n) if not np.isnan(coords[j, 0])), -1)\n                if first_valid >= 0:\n                    for j in range(first_valid - 1, -1, -1):\n                        direction = np.random.normal(0, 1, 3)\n                        direction = direction / (np.linalg.norm(direction) + 1e-10) * typical_step\n                        coords[j] = coords[j + 1] - direction\n                else:\n                    coords = generate_basic_structure_coords(n)\n                    break\n            else:\n                prev_valid = next((j for j in range(i-1, -1, -1) if not np.isnan(coords[j, 0])), -1)\n                if prev_valid >= 0:\n                    direction = np.random.normal(0, 1, 3)\n                    direction = direction / (np.linalg.norm(direction) + 1e-10) * typical_step\n                    coords[i] = coords[prev_valid] + direction\n    \n    # Final cleanup\n    coords = np.nan_to_num(coords)\n    return coords\n\n\ndef generate_basic_structure(sequence):\n    \"\"\"Generate a simple helical structure.\"\"\"\n    return generate_basic_structure_coords(len(sequence))\n\n\ndef generate_basic_structure_coords(n_residues):\n    \"\"\"Generate basic helical coordinates.\"\"\"\n    coords = np.zeros((n_residues, 3))\n    for i in range(n_residues):\n        angle = i * 0.6\n        coords[i] = [10.0 * np.cos(angle), 10.0 * np.sin(angle), i * 2.5]\n    return coords\n\n# Cell 6: RNA Geometric Constraints\ndef adaptive_rna_constraints(coordinates, sequence, confidence=1.0):\n    \"\"\"\n    Apply RNA geometric constraints with adaptive strength based on confidence.\n    \"\"\"\n    refined_coords = coordinates.copy()\n    n_residues = len(sequence)\n    \n    constraint_strength = 0.8 * (1.0 - min(confidence, 0.8))\n    \n    # Sequential distance constraints\n    seq_min_dist, seq_max_dist = 5.5, 6.5\n    \n    for i in range(n_residues - 1):\n        current_pos = refined_coords[i]\n        next_pos = refined_coords[i + 1]\n        current_dist = np.linalg.norm(next_pos - current_pos)\n        \n        if current_dist < seq_min_dist or current_dist > seq_max_dist:\n            target_dist = (seq_min_dist + seq_max_dist) / 2\n            direction = next_pos - current_pos\n            direction = direction / (np.linalg.norm(direction) + 1e-10)\n            adjustment = (target_dist - current_dist) * constraint_strength\n            refined_coords[i + 1] = current_pos + direction * (current_dist + adjustment)\n    \n    # Steric clash prevention\n    min_allowed_distance = 3.8\n    dist_matrix = distance_matrix(refined_coords, refined_coords)\n    severe_clashes = np.where((dist_matrix < min_allowed_distance) & (dist_matrix > 0))\n    \n    for idx in range(len(severe_clashes[0])):\n        i, j = severe_clashes[0][idx], severe_clashes[1][idx]\n        if abs(i - j) <= 1 or i >= j:\n            continue\n        \n        pos_i, pos_j = refined_coords[i], refined_coords[j]\n        current_dist = dist_matrix[i, j]\n        direction = pos_j - pos_i\n        direction = direction / (np.linalg.norm(direction) + 1e-10)\n        adjustment = (min_allowed_distance - current_dist) * constraint_strength\n        refined_coords[i] = pos_i - direction * (adjustment / 2)\n        refined_coords[j] = pos_j + direction * (adjustment / 2)\n    \n    return refined_coords\n\n# Cell 7: De Novo Structure Generation\ndef generate_rna_structure(sequence, seed=None):\n    \"\"\"Generate a more realistic RNA structure prediction.\"\"\"\n    if seed is not None:\n        np.random.seed(seed)\n        random.seed(seed)\n    \n    n_residues = len(sequence)\n    coordinates = np.zeros((n_residues, 3))\n    \n    # Initialize first residues\n    for i in range(min(3, n_residues)):\n        angle = i * 0.6\n        coordinates[i] = [10.0 * np.cos(angle), 10.0 * np.sin(angle), i * 2.5]\n    \n    current_direction = np.array([0.0, 0.0, 1.0])\n    complementary = {'G': 'C', 'C': 'G', 'A': 'U', 'U': 'A'}\n    \n    for i in range(3, n_residues):\n        current_base = sequence[i]\n        has_pair = False\n        pair_idx = -1\n        \n        window_size = min(i, 15)\n        for j in range(i - window_size, i):\n            if j >= 0 and sequence[j] == complementary.get(current_base, 'X'):\n                has_pair = True\n                pair_idx = j\n                break\n        \n        if has_pair and i - pair_idx <= 10 and random.random() < 0.7:\n            pair_pos = coordinates[pair_idx]\n            random_offset = np.random.normal(0, 1, 3) * 2.0\n            base_pair_distance = 10.0 + random.uniform(-1.0, 1.0)\n            \n            center = np.mean(coordinates[:i], axis=0)\n            direction = center - pair_pos\n            direction = direction / (np.linalg.norm(direction) + 1e-10)\n            \n            coordinates[i] = pair_pos + direction * base_pair_distance + random_offset\n            current_direction = np.random.normal(0, 0.3, 3)\n            current_direction = current_direction / (np.linalg.norm(current_direction) + 1e-10)\n        else:\n            if random.random() < 0.3:\n                angle = random.uniform(0.2, 0.6)\n                axis = np.random.normal(0, 1, 3)\n                axis = axis / (np.linalg.norm(axis) + 1e-10)\n                rotation = R.from_rotvec(angle * axis)\n                current_direction = rotation.apply(current_direction)\n            else:\n                current_direction += np.random.normal(0, 0.15, 3)\n                current_direction = current_direction / (np.linalg.norm(current_direction) + 1e-10)\n            \n            step_size = random.uniform(3.5, 4.5)\n            coordinates[i] = coordinates[i - 1] + step_size * current_direction\n    \n    return coordinates\n\n# Cell 8: Main Prediction Function\ndef predict_rna_structures(sequence, target_id, train_seqs_df, train_coords_dict, \n                          n_predictions=5, temporal_cutoff=None):\n    global rep_training_seqs,max_length_diff,min_tbm_target_length,num_tbm_predictions,skip_dot_product;\n    \"\"\"Generate multiple structure predictions for an RNA sequence.\"\"\"\n    predictions = [];\n    n_prediction_buff = n_predictions*3;\n    \n    if len(sequence) < min_tbm_target_length:\n        print(\"Skip TBM for short sequence:\", sequence )    \n    else:\n        # Find similar sequences\n        \n        if num_tbm_predictions - num_eba_predictions > 0:\n            similar_seqs = find_similar_sequences(\n                sequence, train_seqs_df, train_coords_dict,\n                temporal_cutoff=temporal_cutoff, top_n=n_prediction_buff\n            )\n        \n            # Use templates if found\n            if similar_seqs:\n                for template_id, template_seq, similarity, template_coords in similar_seqs:\n                    adapted_coords = adapt_template_to_query(sequence, template_seq, template_coords)\n                    \n                    if adapted_coords is not None:\n                        refined_coords = adaptive_rna_constraints(adapted_coords, sequence, confidence=similarity)\n                        random_scale = max(0.05, 0.8 - similarity)\n                        randomized_coords = refined_coords + np.random.normal(0, random_scale, refined_coords.shape)\n                        predictions.append(randomized_coords)\n                        \n                        if len(predictions) >= n_prediction_buff:\n                            break\n                        \n        # Add EBA hits                \n        if num_eba_predictions > 0:\n            try:\n                slen = len(sequence);\n                eba_predictions = [];\n                eba_predictions2 = [];\n                for dot_flag in [False,True]:\n                    if skip_dot_product and dot_flag:\n                        continue;\n                    if dot_flag:\n                        # To get more diverge hits\n                        res = run_eba(sequence,use_dot_product=dot_flag,num_prefilter_seqs=10000//slen+20);\n                    else:\n                        res = run_eba(sequence,use_dot_product=dot_flag,num_prefilter_seqs=10000//slen+5);\n                    if len(res) > 0:\n                        # Take scores in missing regions into account.\n                        sortedres = sort_matrix_align_hit(res);\n                        if len(sortedres) > 5:\n                            print(\"Top 5 EBA Templates:\",[(x.get(\"name\",\"\"),x.get(\"desc\",\"\"),x[\"score\"]) for x in sortedres[0:5]]);\n                        for ii in range(min(len(sortedres),10)):\n                            mfa = a3m_to_mfa(sequence,sortedres[ii][\"seq\"]);\n                            with open(sortedres[ii][\"json\"]) as fin:\n                                pdat = json.load(fin);\n                            xx = pdat[\"x_1\"];\n                            yy = pdat[\"y_1\"];\n                            zz = pdat[\"z_1\"];\n                            coordinates = [];\n                            for jj in range(len(xx)):\n                                if len(xx[jj]) == 0 or len(yy[jj]) == 0 or len(zz[jj]) == 0:\n                                    coordinates.append([float(\"nan\"),float(\"nan\"),float(\"nan\")]);\n                                else:\n                                    coordinates.append(\n                                        [float(xx[jj]),float(yy[jj]),float(zz[jj])]\n                                    );\n                                    \n                            assert len(pdat[\"all_bases\"]) == len(coordinates);\n                            template_coords = np.array(coordinates,dtype=np.float32);\n                            adapted_coords = adapt_template_to_query(sequence,pdat[\"all_bases\"],template_coords,precomputed_alignment=(mfa[0], mfa[1]))\n                            matchcount = 0;\n                            hitcount = 0;\n                            for p,q in zip(mfa[0],mfa[1]):\n                                if p == \"-\" or q == \"-\":\n                                    continue;\n                                hitcount += 1;\n                                if p == q:\n                                    matchcount += 1;\n                            if hitcount > 0:\n                                similarity = matchcount/float(hitcount);\n                                refined_coords = adaptive_rna_constraints(adapted_coords, sequence, confidence=similarity)\n                                random_scale = max(0.05, 0.8 - similarity)\n                                randomized_coords = refined_coords + np.random.normal(0, random_scale, refined_coords.shape)\n                                if dot_flag:\n                                    eba_predictions2.append(randomized_coords);\n                                else:\n                                    eba_predictions.append(randomized_coords);\n                            \n                    updated_predictions = [];\n                    for _ in range(num_eba_predictions):\n                        if len(eba_predictions) > 0:\n                            updated_predictions.append(eba_predictions.pop(0));\n                        if len(updated_predictions) >= num_eba_predictions:\n                            break;\n                        if len(eba_predictions2) > 0:\n                            updated_predictions.append(eba_predictions2.pop(0));\n                        if len(updated_predictions) >= num_eba_predictions:\n                            break;\n                            \n                    while len(updated_predictions) < num_tbm_predictions:\n                        if len(predictions) > 0:\n                            updated_predictions.append(predictions.pop(0));\n                        else:\n                            break;\n                            \n                    while len(updated_predictions) < n_prediction_buff:\n                        if len(predictions) > 0:\n                            updated_predictions.append(predictions.pop(0));\n                        if len(eba_predictions) > 0 and len(updated_predictions) < n_prediction_buff:\n                            updated_predictions.append(eba_predictions.pop(0));\n                        if len(eba_predictions2) > 0 and len(updated_predictions) < n_prediction_buff:\n                            updated_predictions.append(eba_predictions2.pop(0));\n                        if len(eba_predictions2) == 0 and len(eba_predictions) == 0 and len(predictions) == 0:\n                            break;\n                    predictions = updated_predictions;\n            except:\n                import traceback;\n                traceback.print_exc();\n\n    if len(predictions) > n_predictions:\n        removed_coords = [];\n        retained_coords = [];\n        plen = len(predictions);\n        retained_coords.append(predictions[0]);\n        for i in range(0,plen):\n            remflag = False;\n            for j in range(len(retained_coords)):\n                score = get_tmscore(\n                    retained_coords[j],\n                    predictions[i]\n                );\n                if score > similarity_tmscore_max:\n                    remflag = True;\n                    break;\n                    \n            if remflag:\n                removed_coords.append(predictions[i]);\n            else:\n                retained_coords.append(predictions[i]);\n                \n            if len(retained_coords) >= n_predictions:\n                break;\n                \n        predictions = retained_coords;\n        \n        if len(predictions) < n_predictions:\n            predictions.extend(removed_coords);\n            \n    # Fill remaining with de novo structures\n    while len(predictions) < n_predictions:\n        seed_value = hash(target_id) % 10000 + len(predictions) * 1000\n        de_novo_coords = generate_rna_structure(sequence, seed=seed_value)\n        if False:\n            refined_de_novo = adaptive_rna_constraints(de_novo_coords, sequence, confidence=0.2)\n            predictions.append(refined_de_novo)\n        else:\n            predictions.append(de_novo_coords*0.0); # de novo の効果の確認のためゼロ座標を返す\n    return predictions[:n_predictions]\n\n# Cell 9: Generate Predictions\ndef generate_predictions_for_dataset(seqs_df, train_seqs_df, train_coords_dict, dataset_name=\"dataset\"):\n    \"\"\"Generate predictions for a given sequence dataset.\"\"\"\n    all_predictions = []\n    start_time = time.time()\n    total_targets = len(seqs_df)\n    \n    print(f\"\\n=== Generating Predictions for {dataset_name} ({total_targets} sequences) ===\")\n    \n    for idx, row in seqs_df.iterrows():\n        target_id = row['target_id']\n        sequence = row['sequence']\n        temporal_cutoff = row.get('temporal_cutoff', None)\n        \n        if idx % 5 == 0:\n            elapsed = time.time() - start_time\n            if idx > 0:\n                est_remaining = elapsed / (idx + 1) * (total_targets - idx - 1)\n                print(f\"Processing {idx+1}/{total_targets}: {target_id} ({len(sequence)} nt), \"\n                      f\"elapsed: {elapsed:.1f}s, remaining: {est_remaining:.1f}s\")\n            else:\n                print(f\"Processing {idx+1}/{total_targets}: {target_id} ({len(sequence)} nt)\")\n        \n        predictions = predict_rna_structures(\n            sequence, target_id, train_seqs_df, train_coords_dict,\n            n_predictions=5, temporal_cutoff=temporal_cutoff\n        )\n        \n        for j in range(len(sequence)):\n            pred_row = {\n                'ID': f\"{target_id}_{j+1}\",\n                'resname': sequence[j],\n                'resid': j + 1\n            }\n            for i in range(5):\n                pred_row[f'x_{i+1}'] = predictions[i][j][0]\n                pred_row[f'y_{i+1}'] = predictions[i][j][1]\n                pred_row[f'z_{i+1}'] = predictions[i][j][2]\n            all_predictions.append(pred_row)\n    \n    # Create DataFrame\n    df = pd.DataFrame(all_predictions)\n    column_order = ['ID', 'resname', 'resid']\n    for i in range(1, 6):\n        for coord in ['x', 'y', 'z']:\n            column_order.append(f'{coord}_{i}')\n    df = df[column_order]\n    \n    print(f\"\\nGenerated predictions for {total_targets} sequences\")\n    print(f\"Total runtime: {time.time() - start_time:.1f} seconds\")\n    print(f\"Output shape: {df.shape}\")\n    \n    return df\n\n# Cell 10: Generate Predictions for TEST Set (Submission)\n# Generate test predictions for submission\ntest_pred_df = generate_predictions_for_dataset(\n    test_seqs, train_seqs, train_coords_dict,\n    dataset_name=\"Test\"\n)\ntest_pred_df.head()\n\n# Cell 11: Save Submission\ntest_pred_df.to_csv('submission_tbm.csv', index=False)\nprint(\"\\n=== Submission file saved ===\")\nprint(f\"submission.csv shape: {test_pred_df.shape}\")\nprint(f\"\\nFirst few rows:\")\ntest_pred_df.head(10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T06:06:48.303348Z","iopub.execute_input":"2026-03-14T06:06:48.303624Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub_tbm = pd.read_csv(\"/kaggle/working/submission_tbm.csv\")\nsub_tbm","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pwd","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"continue RNAPro\n---  \n\n","metadata":{}},{"cell_type":"code","source":"cd RNAPro","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python preprocess/convert_templates_to_pt_files.py --input_csv /kaggle/working/submission_tbm.csv --output_name templates.pt","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DIST = \"/kaggle/working/RNAPro/release_data/ccd_cache/\"\n!mkdir -p $DIST","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Recomputed the cdd cache using python preprocess/gen_ccd_cache.py\n# Exported as dataset here: https://www.kaggle.com/datasets/jaejohn/rnapro-ccd-cache\n\n# You will need the following packages to recompute it on our own\n\n# pip install gemmi\n# pip install pdbeccdutils","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# updated file paths\n!cp /kaggle/input/rnapro-ccd-cache/ccd_cache/components.cif $DIST\n!cp /kaggle/input/rnapro-ccd-cache/ccd_cache/components.cif.rdkit_mol.pkl $DIST","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Inference  \n---","metadata":{}},{"cell_type":"code","source":"# %%python\nimport pandas as pd\ndf = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\")\nif not IS_SCORING_RUN:\n    df = df.head(5)\ndf.to_csv('/kaggle/working/sample_sequences.csv', index=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile runner/inference.py\n# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.\n# SPDX-License-Identifier: Apache-2.0\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n#     http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport os\nimport shutil\nimport logging\nimport traceback\nimport warnings\nimport argparse\nfrom contextlib import nullcontext\nfrom os.path import join as opjoin\nfrom typing import Any, Mapping\n\nimport json\nimport torch\nimport pandas as pd\nimport numpy as np\nfrom biotite.structure.io import pdbx\n\nfrom configs.configs_base import configs as configs_base\nfrom configs.configs_data import data_configs\nfrom configs.configs_inference import inference_configs\nfrom runner.dumper import DataDumper\n\nfrom rnapro.config import parse_sys_args\nfrom rnapro.config.config import ConfigManager, ArgumentNotSet\nfrom rnapro.data.infer_data_pipeline import get_inference_dataloader\nfrom rnapro.model.RNAPro import RNAPro\nfrom rnapro.utils.distributed import DIST_WRAPPER\nfrom rnapro.utils.seed import seed_everything\nfrom rnapro.utils.torch_utils import to_device\n\nwarnings.filterwarnings(\"ignore\", category=UserWarning)\nwarnings.filterwarnings(\"ignore\", category=FutureWarning)\nwarnings.filterwarnings(\"ignore\", category=DeprecationWarning)\n\nlogger = logging.getLogger(__name__)\n\n# Silence all info logging\nlogging.basicConfig(level=logging.WARNING)\n# Silence dataloader logging specifically\nlogging.getLogger(\"rnapro.data\").setLevel(logging.WARNING)\nlogging.getLogger(\"rnapro\").setLevel(logging.WARNING)\n\n\ndef parse_configs(\n    configs: dict, arg_str: str = None, fill_required_with_null: bool = False\n):\n    \"\"\"\n    Parses and merges configuration settings from a dictionary and command-line arguments.\n\n    Args:\n        configs (dict): A dictionary containing initial configuration settings.\n        arg_str (str, optional): A string representing command-line arguments. Defaults to None.\n        fill_required_with_null (bool, optional):\n            A boolean flag indicating whether required values should be filled with `None` if not provided. Defaults to False.\n\n    Returns:\n        ConfigDict: The merged configuration dictionary.\n    \"\"\"\n    manager = ConfigManager(configs, fill_required_with_null=fill_required_with_null)\n    parser = argparse.ArgumentParser()\n\n    # This is new\n    parser.add_argument(\n        \"--max_len\",\n        type=int,\n        default=10000,\n        required=False,\n        help=\"Maximum length of the sequence. Longer sequences will be skipped during inference\"\n    )\n\n    # Register arguments\n    for key, (\n        dtype,\n        default_value,\n        allow_none,\n        required,\n    ) in manager.config_infos.items():\n        # All config use str type, strings will be converted to real dtype later\n        parser.add_argument(\n            \"--\" + key, type=str, default=ArgumentNotSet(), required=required\n        )\n    # Merge user commandline pargs with default ones\n    merged_configs = manager.merge_configs(\n        vars(parser.parse_args(arg_str.split())) if arg_str else {}\n    )\n\n    max_len = parser.parse_args(arg_str.split()).max_len\n    merged_configs.max_len = max_len\n\n    return merged_configs\n\n\nclass dotdict(dict):\n    __setattr__ = dict.__setitem__\n    __delattr__ = dict.__delitem__\n\n    def __getattr__(self, name):\n        try:\n            return self[name]\n        except KeyError:\n            raise AttributeError(name)\n\n\nclass InferenceRunner(object):\n    def __init__(self, configs: Any) -> None:\n        self.configs = configs\n        self.init_env()\n        self.init_basics()\n        self.init_model()\n        self.load_checkpoint()\n        self.init_dumper(\n            need_atom_confidence=configs.need_atom_confidence,\n            sorted_by_ranking_score=configs.sorted_by_ranking_score,\n        )\n\n    def init_env(self) -> None:\n        self.print(\n            f\"Distributed environment: world size: {DIST_WRAPPER.world_size}, \"\n            + f\"global rank: {DIST_WRAPPER.rank}, local rank: {DIST_WRAPPER.local_rank}\"\n        )\n        self.use_cuda = torch.cuda.device_count() > 0\n        if self.use_cuda:\n            self.device = torch.device(\"cuda:{}\".format(DIST_WRAPPER.local_rank))\n            os.environ[\"CUDA_DEVICE_ORDER\"] = \"PCI_BUS_ID\"\n            all_gpu_ids = \",\".join(str(x) for x in range(torch.cuda.device_count()))\n            devices = os.getenv(\"CUDA_VISIBLE_DEVICES\", all_gpu_ids)\n            logging.info(\n                f\"LOCAL_RANK: {DIST_WRAPPER.local_rank} - CUDA_VISIBLE_DEVICES: [{devices}]\"\n            )\n            torch.cuda.set_device(self.device)\n        else:\n            self.device = torch.device(\"cpu\")\n        if self.configs.use_deepspeed_evo_attention:\n            env = os.getenv(\"CUTLASS_PATH\", None)\n            self.print(f\"env: {env}\")\n            assert (\n                env is not None\n            ), \"if use ds4sci, set `CUTLASS_PATH` env as https://www.deepspeed.ai/tutorials/ds4sci_evoformerattention/\"\n            if env is not None:\n                logging.info(\n                    \"The kernels will be compiled when DS4Sci_EvoformerAttention is called for the first time.\"\n                )\n        use_fastlayernorm = os.getenv(\"LAYERNORM_TYPE\", None)\n        if use_fastlayernorm == \"fast_layernorm\":\n            logging.info(\n                \"The kernels will be compiled when fast_layernorm is called for the first time.\"\n            )\n\n        logging.info(\"Finished init ENV.\")\n\n    def init_basics(self) -> None:\n        self.dump_dir = self.configs.dump_dir\n        self.error_dir = opjoin(self.dump_dir, \"ERR\")\n        os.makedirs(self.dump_dir, exist_ok=True)\n        os.makedirs(self.error_dir, exist_ok=True)\n\n    def init_model(self) -> None:\n        self.model = RNAPro(self.configs).to(self.device)\n        num_params = sum(p.numel() for p in self.model.parameters())\n        self.print(f\"Total number of parameters: {num_params:,}\")\n\n    def load_checkpoint(self) -> None:\n        checkpoint_path = self.configs.load_checkpoint_path\n\n        if not os.path.exists(checkpoint_path):\n            raise Exception(f\"Given checkpoint path not exist [{checkpoint_path}]\")\n        self.print(\n            f\"Loading from {checkpoint_path}, strict: {self.configs.load_strict}\"\n        )\n        checkpoint = torch.load(checkpoint_path, self.device)\n\n        sample_key = [k for k in checkpoint[\"model\"].keys()][0]\n        # self.print(f\"Sampled key: {sample_key}\")\n        if sample_key.startswith(\"module.\"):  # DDP checkpoint has module. prefix\n            checkpoint[\"model\"] = {\n                k[len(\"module.\"):]: v for k, v in checkpoint[\"model\"].items()\n            }\n        self.model.load_state_dict(\n            state_dict=checkpoint[\"model\"],\n            strict=True,\n        )\n        self.model.eval()\n        # self.print(\"Finish loading checkpoint.\")\n\n    def init_dumper(\n        self, need_atom_confidence: bool = False, sorted_by_ranking_score: bool = True\n    ):\n        self.dumper = DataDumper(\n            base_dir=self.dump_dir,\n            need_atom_confidence=need_atom_confidence,\n            sorted_by_ranking_score=sorted_by_ranking_score,\n        )\n\n    def print_dict(self, d):\n        for k, v in d.items():\n            if isinstance(v, torch.Tensor):\n                print(f\"{k}: \", v.shape)\n            else:\n                pass\n                # print(f\"{k}: {v}\")\n\n    # Adapted from runner.train.Trainer.evaluate\n    @torch.no_grad()\n    def predict(self, data: Mapping[str, Mapping[str, Any]]) -> dict[str, torch.Tensor]:\n        eval_precision = {\n            \"fp32\": torch.float32,\n            \"bf16\": torch.bfloat16,\n            \"fp16\": torch.float16,\n        }[self.configs.dtype]\n        # print(\"eval_precision: \", eval_precision)\n        enable_amp = (\n            torch.autocast(device_type=\"cuda\", dtype=eval_precision)\n            if torch.cuda.is_available()\n            else nullcontext()\n        )\n        #         print('input_feature_dict: ', self.print_dict(data[\"input_feature_dict\"]))\n        #         exit(0)\n\n        data = to_device(data, self.device)\n        with enable_amp:\n            prediction, _, _ = self.model(\n                input_feature_dict=data[\"input_feature_dict\"],\n                label_full_dict=None,\n                label_dict=None,\n                mode=\"inference\",\n            )\n\n        return prediction\n\n    def print(self, msg: str):\n        if DIST_WRAPPER.rank == 0:\n            # logger.info(msg)\n            print(msg)\n\n    def update_model_configs(self, new_configs: Any) -> None:\n        self.model.configs = new_configs\n\n\ndef update_inference_configs(configs: Any, N_token: int):\n    # Setting the default inference configs for different N_token and N_atom\n    # when N_token is larger than 3000, the default config might OOM even on a\n    # A100 80G GPUS,\n    if N_token > 3840:\n        configs.skip_amp.confidence_head = False\n        configs.skip_amp.sample_diffusion = False\n    elif N_token > 2560:\n        configs.skip_amp.confidence_head = False\n        configs.skip_amp.sample_diffusion = True\n    else:\n        configs.skip_amp.confidence_head = True\n        configs.skip_amp.sample_diffusion = True\n    return configs\n\n\ndef infer_predict(runner: InferenceRunner, configs: Any) -> None:\n    # Data\n    # logger.info(f\"Loading data from {configs.input_json_path}\")\n    try:\n        dataloader = get_inference_dataloader(configs=configs)\n    except Exception as e:\n        error_message = f\"{e}:\\n{traceback.format_exc()}\"\n        logger.info(error_message)\n        with open(opjoin(runner.error_dir, \"error.txt\"), \"a\") as f:\n            f.write(error_message)\n        return\n\n    num_data = len(dataloader.dataset)\n    for seed in configs.seeds:\n        seed_everything(seed=seed, deterministic=configs.deterministic)\n        for batch in dataloader:\n            try:\n                data, atom_array, data_error_message = batch[0]\n                sample_name = data[\"sample_name\"]\n\n                if len(data_error_message) > 0:\n                    logger.info(data_error_message)\n                    with open(opjoin(runner.error_dir, f\"{sample_name}.txt\"), \"a\") as f:\n                        f.write(data_error_message)\n                    continue\n\n                logger.info(\n                    (\n                        f\"[Rank {DIST_WRAPPER.rank} ({data['sample_index'] + 1}/{num_data})] {sample_name}: \"\n                        f\"N_asym {data['N_asym'].item()}, N_token {data['N_token'].item()}, \"\n                        f\"N_atom {data['N_atom'].item()}, N_msa {data['N_msa'].item()}\"\n                    )\n                )\n                new_configs = update_inference_configs(configs, data[\"N_token\"].item())\n                runner.update_model_configs(new_configs)\n                prediction = runner.predict(data)\n                runner.dumper.dump(\n                    dataset_name=\"\",\n                    pdb_id=sample_name,\n                    seed=seed,\n                    pred_dict=prediction,\n                    atom_array=atom_array,\n                    entity_poly_type=data[\"entity_poly_type\"],\n                )\n\n                logger.info(\n                    f\"[Rank {DIST_WRAPPER.rank}] {data['sample_name']} succeeded - \"\n                    f\"Results saved to {configs.dump_dir}\"\n                )\n                torch.cuda.empty_cache()\n            except Exception as e:\n                error_message = f\"[Rank {DIST_WRAPPER.rank}]{data['sample_name']} {e}:\\n{traceback.format_exc()}\"\n                logger.info(error_message)\n                # Save error info\n                with open(opjoin(runner.error_dir, f\"{sample_name}.txt\"), \"a\") as f:\n                    f.write(error_message)\n                if hasattr(torch.cuda, \"empty_cache\"):\n                    torch.cuda.empty_cache()\n\n\n# data helper\ndef make_dummy_solution(valid_df):\n    solution = dotdict()\n    for i, row in valid_df.iterrows():\n        target_id = row.target_id\n        sequence = row.sequence\n        solution[target_id] = dotdict(\n            target_id=target_id,\n            sequence=sequence,\n            coord=[],\n        )\n    return solution\n\n\ndef solution_to_submit_df(solution):\n    submit_df = []\n    for k, s in solution.items():\n        df = coord_to_df(s.sequence, s.coord, s.target_id)\n        submit_df.append(df)\n\n    submit_df = pd.concat(submit_df)\n    return submit_df\n\n\ndef coord_to_df(sequence, coord, target_id):\n    L = len(sequence)\n    df = pd.DataFrame()\n    df[\"ID\"] = [f\"{target_id}_{i + 1}\" for i in range(L)]\n    df[\"resname\"] = [s for s in sequence]\n    df[\"resid\"] = [i + 1 for i in range(L)]\n\n    num_coord = len(coord)\n    for j in range(num_coord):\n        df[f\"x_{j+1}\"] = coord[j][:, 0]\n        df[f\"y_{j+1}\"] = coord[j][:, 1]\n        df[f\"z_{j+1}\"] = coord[j][:, 2]\n    return df\n\n\ndef main(configs: Any) -> None:\n    # Runner\n    runner = InferenceRunner(configs)\n    infer_predict(runner, configs)\n\n\ndef create_input_json(sequence, target_id):\n    # print(\"input_no_msa\")\n    input_json = [\n        {\n            \"sequences\": [\n                {\n                    \"rnaSequence\": {\n                        \"sequence\": sequence,\n                        \"count\": 1,\n                    }\n                }\n            ],\n            \"name\": target_id,\n        }\n    ]\n    return input_json\n\n\ndef extract_c1_coordinates(cif_file_path):\n    try:\n        # Read the CIF file using the correct biotite method\n        with open(cif_file_path, \"r\") as f:\n            cif_data = pdbx.CIFFile.read(f)\n\n        # Get structure from CIF data\n        atom_array = pdbx.get_structure(cif_data, model=1)\n\n        # Clean atom names and find C1' atoms\n        atom_names_clean = np.char.strip(atom_array.atom_name.astype(str))\n        mask_c1 = atom_names_clean == \"C1'\"\n        c1_atoms = atom_array[mask_c1]\n\n        if len(c1_atoms) == 0:\n            print(f\"Warning: No C1' atoms found in {cif_file_path}\")\n            return None\n\n        # Sort by residue ID and return coordinates\n        sort_indices = np.argsort(c1_atoms.res_id)\n        c1_atoms_sorted = c1_atoms[sort_indices]\n        c1_coords = c1_atoms_sorted.coord\n\n        return c1_coords\n    except Exception as e:\n        print(f\"Error extracting C1' coordinates from {cif_file_path}: {e}\")\n        return None\n\n\ndef process_sequence(sequence, target_id, temp_dir):\n    # Create input JSON\n    input_json = create_input_json(sequence, target_id)\n\n    # Save JSON to temporary file\n    os.makedirs(temp_dir, exist_ok=True)\n    input_json_path = os.path.join(temp_dir, f\"{target_id}_input.json\")\n    with open(input_json_path, \"w\") as f:\n        json.dump(input_json, f, indent=4)\n\n\ndef run_ptx(target_id, sequence, configs, solution, template_idx, runner):\n    # Create directories\n    temp_dir = f\"./{configs.dump_dir}/input\"  # Same as in kaggle_inference.py\n    output_dir = f\"./{configs.dump_dir}/output\"  # Same as in kaggle_inference.py\n    os.makedirs(temp_dir, exist_ok=True)\n    os.makedirs(output_dir, exist_ok=True)\n\n    process_sequence(sequence=sequence, target_id=target_id, temp_dir=temp_dir)\n    configs.input_json_path = os.path.join(temp_dir, f\"{target_id}_input.json\")\n    configs.template_idx = int(template_idx)\n\n    infer_predict(runner, configs)\n\n    cif_file_path = (\n        f\"{configs.dump_dir}/{target_id}/seed_42/predictions/{target_id}_sample_0.cif\"\n    )\n    cif_new_path = f\"{configs.dump_dir}/{target_id}/seed_42/predictions/{target_id}_sample_{template_idx}_new.cif\"\n    shutil.copy(cif_file_path, cif_new_path)\n    coord = extract_c1_coordinates(cif_file_path)\n    if coord is None:\n        coord = np.zeros((len(sequence), 3), dtype=np.float32)\n    elif coord.shape[0] < (len(sequence)):\n        pad_len = len(sequence) - coord.shape[0]\n        pad = np.zeros((pad_len, 3), dtype=np.float32)\n        coord = np.concatenate([coord, pad], axis=0)\n    solution[target_id].coord.append(coord)\n\n\ndef run() -> None:\n    LOG_FORMAT = \"%(asctime)s,%(msecs)-3d %(levelname)-8s [%(filename)s:%(lineno)s %(funcName)s] %(message)s\"\n    logging.basicConfig(\n        format=LOG_FORMAT,\n        level=logging.WARNING,\n        datefmt=\"%Y-%m-%d %H:%M:%S\",\n        filemode=\"w\",\n    )\n    # Silence dataloader and rnapro module logging\n    logging.getLogger(\"rnapro.data\").setLevel(logging.WARNING)\n    logging.getLogger(\"rnapro\").setLevel(logging.WARNING)\n    configs_base[\"use_deepspeed_evo_attention\"] = (\n        os.environ.get(\"USE_DEEPSPEED_EVO_ATTENTION\", False) == \"true\"\n    )\n    configs = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n    configs = parse_configs(\n        configs=configs,\n        arg_str=parse_sys_args(),\n        fill_required_with_null=True,\n    )\n\n    valid_df = pd.read_csv(configs.sequences_csv)\n    print(f\"\\n -> Loaded {len(valid_df)} sequence(s)\")\n\n    # Build model and load checkpoint once before looping over sequences\n\n    print('\\n -> Building model and loading checkpoint')\n    runner = InferenceRunner(configs)\n    print('\\n -> Done, starting inference...')\n\n    solution = make_dummy_solution(valid_df)\n    for idx, row in valid_df.iterrows():\n        print(f\"\\n -> Sequence {row.target_id}: {row.sequence}\")\n\n        if len(row.sequence) > configs.max_len:\n            print(f'Sequence is too long ({len(row.sequence)} > {configs.max_len}), skipping')\n            for template_idx in range(5):\n                coord = np.zeros((len(row.sequence), 3), dtype=np.float32)\n                solution[row.target_id].coord.append(coord)\n            continue\n\n        try:\n            target_id = row.target_id\n            sequence = row.sequence\n            for template_idx in range(5):\n                print()\n                run_ptx(\n                    target_id=target_id,\n                    sequence=sequence,\n                    configs=configs,\n                    solution=solution,\n                    template_idx=template_idx,\n                    runner=runner,\n                )\n        except Exception as e:\n            print(f\"Error processing {row.target_id}: {e}\")\n            continue\n\n    print('\\n\\n -> Inference done ! Saving to submission.csv')\n    submit_df = solution_to_submit_df(solution)\n    submit_df = submit_df.fillna(0.0)\n    submit_df.to_csv(\"./submission.csv\", index=False)\n\n\nif __name__ == \"__main__\":\n    run()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile rnapro_inference_kaggle.sh\n\nexport LAYERNORM_TYPE=torch # fast_layernorm, torch\n\n\n# Inference parameters (RNAPro)\nSEED=42\nN_SAMPLE=1\nN_STEP=200\nN_CYCLE=10\n\n# Paths\nDUMP_DIR=\"../output\"\n# Set a valid checkpoint file path below\nCHECKPOINT_PATH=\"../rnapro-private-best-500m.ckpt\"\n\n# Template/MSA settings\nTEMPLATE_DATA=\"./release_data/kaggle/templates.pt\"\n# Note: template_idx supports 5 choices and maps to top-k:\n# 0->top1, 1->top2, 2->top3, 3->top4, 4->top5\nTEMPLATE_IDX=0\nRNA_MSA_DIR=\"/kaggle/input/stanford-rna-3d-folding-2/MSA\"\n\n# SEQUENCES_CSV=\"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\"\nSEQUENCES_CSV=\"/kaggle/working/sample_sequences.csv\"\n\n# RibonanzaNet2 path (keep as-is per request)\nRIBONANZA_PATH=\"/kaggle/input/ribonanzanet2/pytorch/alpha/1/\"\n\n# Model selection: keep to an existing key to align defaults (N_step=200, N_cycle=10)\nMODEL_NAME=\"rnapro_base\"\nmkdir -p \"${DUMP_DIR}\"\n\npython3 runner/inference.py \\\n    --model_name \"${MODEL_NAME}\" \\\n    --seeds ${SEED} \\\n    --dump_dir \"${DUMP_DIR}\" \\\n    --load_checkpoint_path \"${CHECKPOINT_PATH}\" \\\n    --use_msa true \\\n    --use_template \"ca_precomputed\" \\\n    --model.use_template \"ca_precomputed\" \\\n    --model.use_RibonanzaNet2 true \\\n    --model.template_embedder.n_blocks 2 \\\n    --model.ribonanza_net_path \"${RIBONANZA_PATH}\" \\\n    --template_data \"${TEMPLATE_DATA}\" \\\n    --template_idx ${TEMPLATE_IDX} \\\n    --rna_msa_dir \"${RNA_MSA_DIR}\" \\\n    --model.N_cycle ${N_CYCLE} \\\n    --sample_diffusion.N_sample ${N_SAMPLE} \\\n    --sample_diffusion.N_step ${N_STEP} \\\n    --load_strict true \\\n    --num_workers 0 \\\n    --triangle_attention \"torch\" \\\n    --triangle_multiplicative \"torch\" \\\n    --sequences_csv \"${SEQUENCES_CSV}\" \\\n    --max_len 1000\n\n\n# --triangle_attention supports 'triattention', 'cuequivariance', 'deepspeed', 'torch'\n# --triangle_multiplicative supports 'cuequivariance', 'torch'\n# --max_len 1000: Sequences longer than max_len will be skipped to avoid oom","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if not skip_rnapro:\n    os.system(\"bash ./rnapro_inference_kaggle.sh\");\n    os.system(\"mv submission.csv ../submission_rnapro.csv\");","metadata":{"trusted":true,"scrolled":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cd ..","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import csv;\ndef baseid_to_pred(csvfile):\n    ret = {};\n    with open(csvfile, 'r') as fin:\n        reader = csv.DictReader(fin);\n        for row in reader:\n            baseid = row[\"ID\"].split(\"_\")[0];\n            if baseid not in ret:\n                ret[baseid] = {};\n                ret[baseid][\"ids\"] = [];\n                ret[baseid][\"id_to_row\"] = {};\n            ret[baseid][\"ids\"].append(row[\"ID\"]);\n            ret[baseid][\"id_to_row\"][row[\"ID\"]] = row;\n    return ret;\ndef replace_coord(src_rows,src_index,dest_rows,dest_index):\n    updated = False;\n    for rk in list(sorted(src_rows.keys())):\n        if rk not in dest_rows:\n            sys.stderr.write(rk+\" not in dest_rows.\\n\");\n        if src_rows[rk][\"x_\"+str(src_index)] != src_rows[rk][\"y_\"+str(src_index)] or src_rows[rk][\"x_\"+str(src_index)] != src_rows[rk][\"z_\"+str(src_index)]:\n            updated = True;\n        for c in [\"x\",\"y\",\"z\"]:\n            dest_rows[rk][c+\"_\"+str(dest_index)] = src_rows[rk][c+\"_\"+str(src_index)];\n    return updated;","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Merge submission csvs\nbase_data = baseid_to_pred(\"/kaggle/input/stanford-rna-3d-folding-2/sample_submission.csv\");\nwith open(\"/kaggle/input/stanford-rna-3d-folding-2/sample_submission.csv\") as fin:\n    reader = csv.reader(fin);\n    header = next(reader);\n\ntbm_data = baseid_to_pred(\"submission_tbm.csv\");\n\nbaseid_to_modelcount = {};\nfor baseid in list(sorted(base_data.keys())):\n    c1 = base_data[baseid][\"id_to_row\"];\n    baseid_to_modelcount[baseid] = 0;\n    if baseid not in tbm_data:\n        continue;\n    # Because the competition host said that short RNAs are not TBM target.\n    if len(base_data[baseid][\"ids\"]) < min_tbm_target_length:\n        continue;\n    c2 = tbm_data[baseid][\"id_to_row\"];\n    for i in range(1,6):\n        flag = replace_coord(c2,i,c1,i);\n        if flag:\n            if i <= num_tbm_predictions:\n                baseid_to_modelcount[baseid] = i;\nif os.path.exists(\"submission_rnapro.csv\"):\n    rnapro_data = baseid_to_pred(\"submission_rnapro.csv\");\n    \n    for baseid in list(sorted(base_data.keys())):\n        if baseid not in rnapro_data:\n            continue;\n        c1 = base_data[baseid][\"id_to_row\"];\n        c2 = rnapro_data[baseid][\"id_to_row\"];\n        for i in range(baseid_to_modelcount[baseid]+1,6):\n            replace_coord(c2,i,c1,i);\n        \nwith open(\"submission.csv\",\"wt\") as fout:\n    fout.write(\",\".join(header)+\"\\n\");\n    for baseid in list(sorted(base_data.keys())):\n        ids = base_data[baseid][\"ids\"];\n        for k in ids:\n            r = base_data[baseid][\"id_to_row\"][k];\n            line = [];\n            for h in header:\n                line.append(r[h]);\n            fout.write(\",\".join(line)+\"\\n\");","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!head submission.csv","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nsub = pd.read_csv(\"/kaggle/working/submission.csv\")\nsub","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Evaluation  \n---","metadata":{}},{"cell_type":"code","source":"%%writefile /kaggle/working/metric.py\n#!/usr/bin/env python\n# coding: utf-8\n\n# In[ ]:\n\n\nimport os\nimport re\nimport math\nimport pandas as pd\nfrom pathlib import Path\nimport shutil\nimport sys\nimport csv\n\n# ---------------------\n# Helper: parse USalign output\n# ---------------------\ndef parse_tmscore_output(output: str) -> float:\n    matches = re.findall(r'TM-score=\\s+([\\d.]+)', output)\n    if len(matches) < 2:\n        raise ValueError('No TM score found in USalign output')\n    return float(matches[1])\n\n# ---------------------\n# PDB writers\n# ---------------------\n\ndef sanitize(xyz):\n    MIN_COORD=-999.999\n    MAX_COORD=9999.999\n    return min(max(xyz,MIN_COORD),MAX_COORD)\n\ndef write_target_line(atom_name, atom_serial, residue_name, chain_id, residue_num,\n                      x_coord, y_coord, z_coord, occupancy=1.0, b_factor=0.0, atom_type='P') -> str:\n    return f'ATOM  {atom_serial:>5d}  {atom_name:4s}{residue_name:>3s} {chain_id:1s}{residue_num:>4d}    {sanitize(x_coord):>8.3f}{sanitize(y_coord):>8.3f}{sanitize(z_coord):>8.3f}{occupancy:>6.2f}{b_factor:>6.2f}           {atom_type}\\n'\n    \ndef write2pdb(df: pd.DataFrame, xyz_id: int, target_path: str) -> int:\n    \"\"\"\n    Write single-chain PDB (chain 'A') using row['resid'] as residue_num.\n    Raises exceptions on invalid data.\n    \"\"\"\n    resolved_cnt = 0\n    with open(target_path, 'w') as fh:\n        for _, row in df.iterrows():\n            x = row[f'x_{xyz_id}']\n            y = row[f'y_{xyz_id}']\n            z = row[f'z_{xyz_id}']\n            if x > -1e6 and y > -1e6 and z > -1e6:\n                resolved_cnt += 1\n                resid_num = int(row['resid'])\n                fh.write(write_target_line(\"C1'\", resid_num, row['resname'], 'A', resid_num, x, y, z, atom_type='C'))\n    return resolved_cnt\n\ndef write2pdb_singlechain_native(df_native: pd.DataFrame, xyz_id: int, target_path: str) -> int:\n    \"\"\"\n    Write native single-chain using row['resid'] as residue numbers.\n    Assumes all required columns exist and are valid.\n    \"\"\"\n    df_sorted = df_native.copy()\n    df_sorted['__resid_int'] = df_sorted['resid'].astype(int)\n    df_sorted = df_sorted.sort_values('__resid_int').reset_index(drop=True)\n\n    resolved_cnt = 0\n    with open(target_path, 'w') as fh:\n        for _, row in df_sorted.iterrows():\n            x = row[f'x_{xyz_id}']\n            y = row[f'y_{xyz_id}']\n            z = row[f'z_{xyz_id}']\n            if x > -1e6 and y > -1e6 and z > -1e6:\n                resolved_cnt += 1\n                resid_num = int(row['resid'])\n                fh.write(write_target_line(\"C1'\", resid_num, row['resname'], 'A', resid_num, x, y, z, atom_type='C'))\n    return resolved_cnt\n\ndef write2pdb_multichain_from_solution(df_solution: pd.DataFrame, xyz_id: int, target_path: str) -> int:\n    \"\"\"\n    Write multi-chain PDB for native solution using columns 'chain' and 'copy' to assign chain letters.\n    Expects 'resid' convertible to int and chain/copy present. No fallbacks.\n    \"\"\"\n    df_sorted = df_solution.copy()\n    df_sorted['__resid_int'] = df_sorted['resid'].astype(int)\n    df_sorted = df_sorted.sort_values('__resid_int')\n\n    chain_map = {}\n    next_ord = ord('A')\n    written = 0\n    with open(target_path, 'w') as fh:\n        for _, row in df_sorted.iterrows():\n            x = row[f'x_{xyz_id}']\n            y = row[f'y_{xyz_id}']\n            z = row[f'z_{xyz_id}']\n            if not (x > -1e6 and y > -1e6 and z > -1e6):\n                continue\n            chain_val = row['chain']\n            copy_key = int(row['copy'])\n            g = (str(chain_val), copy_key)\n            if g not in chain_map:\n                if next_ord <= ord('Z'):\n                    ch = chr(next_ord)\n                else:\n                    ov = next_ord - ord('Z') - 1\n                    if ov < 26:\n                        ch = chr(ord('a') + ov)\n                    else:\n                        ch = chr(ord('0') + (ov - 26) % 10)\n                chain_map[g] = ch\n                next_ord += 1\n            chain_id = chain_map[g]\n            written += 1\n            resid_num = int(row['resid'])\n            fh.write(write_target_line(\"C1'\", resid_num, row['resname'], chain_id, resid_num, x, y, z, atom_type='C'))\n    return written\n\ndef write2pdb_multichain_from_groups(df_pred: pd.DataFrame, xyz_id: int, target_path: str, groups_list) -> (int, list):\n    \"\"\"\n    Write predicted multichain PDB based on a positional groups_list (tuple per residue: (chain, copy)).\n    Requires groups_list length == number of residues in df_pred (after sorting).\n    Returns (written_count, chain_letters_per_res).\n    \"\"\"\n    df_sorted = df_pred.copy()\n    df_sorted['__resid_int'] = df_sorted['resid'].astype(int)\n    df_sorted = df_sorted.sort_values('__resid_int').reset_index(drop=True)\n\n    if groups_list is None or len(groups_list) != len(df_sorted):\n        raise ValueError(\"groups_list must be provided and match number of residues in predicted df\")\n\n    chain_map = {}\n    next_ord = ord('A')\n    chain_letters = []\n    written = 0\n    with open(target_path, 'w') as fh:\n        for idx, row in df_sorted.iterrows():\n            g = groups_list[idx]\n            if isinstance(g, tuple):\n                gkey = (str(g[0]), int(g[1]))\n            else:\n                gkey = (str(g), None)\n            if gkey not in chain_map:\n                if next_ord <= ord('Z'):\n                    ch = chr(next_ord)\n                else:\n                    ov = next_ord - ord('Z') - 1\n                    if ov < 26:\n                        ch = chr(ord('a') + ov)\n                    else:\n                        ch = chr(ord('0') + (ov - 26) % 10)\n                chain_map[gkey] = ch\n                next_ord += 1\n            chain_id = chain_map[gkey]\n            chain_letters.append(chain_id)\n            x = row[f'x_{xyz_id}']\n            y = row[f'y_{xyz_id}']\n            z = row[f'z_{xyz_id}']\n            if x > -1e6 and y > -1e6 and z > -1e6:\n                written += 1\n                resid_num = int(row['resid'])\n                fh.write(write_target_line(\"C1'\", resid_num, row['resname'], chain_id, resid_num, x, y, z, atom_type='C'))\n    return written, chain_letters\n\ndef write2pdb_singlechain_permuted_pred(df_pred: pd.DataFrame, xyz_id: int, permuted_indices: list, target_path: str) -> int:\n    \"\"\"\n    Create single-chain PDB by concatenating predicted residues in permuted_indices order.\n    Output residue numbers are sequential starting at 1 and increase for every permuted position.\n    Raises exception if indices out of range.\n    \"\"\"\n    df_sorted = df_pred.copy()\n    df_sorted['__resid_int'] = df_sorted['resid'].astype(int)\n    df_sorted = df_sorted.sort_values('__resid_int').reset_index(drop=True)\n\n    written = 0\n    next_res = 1\n    with open(target_path, 'w') as fh:\n        for idx in permuted_indices:\n            if idx < 0 or idx >= len(df_sorted):\n                # strict behavior: raise error for invalid index\n                raise IndexError(f\"permuted index {idx} out of range for predicted residues\")\n            row = df_sorted.iloc[idx]\n            x = row[f'x_{xyz_id}']\n            y = row[f'y_{xyz_id}']\n            z = row[f'z_{xyz_id}']\n            out_resnum = next_res\n            if x > -1e6 and y > -1e6 and z > -1e6:\n                written += 1\n                fh.write(write_target_line(\"C1'\", out_resnum, row['resname'], 'A', out_resnum, x, y, z, atom_type='C'))\n            next_res += 1\n    return written\n\n# ---------------------\n# USalign wrappers\n# ---------------------\ndef run_usalign_raw(predicted_pdb: str, native_pdb: str, usalign_bin='USalign', align_sequence=False, tmscore=None) -> str:\n    cmd = f'{usalign_bin} {predicted_pdb} {native_pdb} -atom \" C1\\'\"'\n    if tmscore is not None:\n        cmd += f' -TMscore {tmscore}'\n        if int(tmscore) == 0:\n            cmd += ' -mm 1 -ter 0'\n    elif not align_sequence:\n        cmd += ' -TMscore 1'\n    return os.popen(cmd).read()\n\ndef parse_usalign_chain_orders(output: str):\n    \"\"\"\n    Parse USalign output for both Structure_1 and Structure_2 chain lists.\n    Returns (chain_list_structure1, chain_list_structure2).\n    Raises if parsing fails to find either line.\n    \"\"\"\n    chain1 = None\n    chain2 = None\n    for line in output.splitlines():\n        line = line.strip()\n        if line.startswith('Name of Structure_1:'):\n            parts = line.split(':')\n            clist = []\n            for part in parts[2:]:\n                token = part.strip()\n                if token == '':\n                    continue\n                token0 = token.split()[0]\n                last = token0.split(',')[-1]\n                ch = re.sub(r'[^A-Za-z0-9]', '', last)\n                if ch:\n                    clist.append(ch)\n            chain1 = clist\n        elif line.startswith('Name of Structure_2:'):\n            parts = line.split(':')\n            clist = []\n            for part in parts[2:]:\n                token = part.strip()\n                if token == '':\n                    continue\n                token0 = token.split()[0]\n                last = token0.split(',')[-1]\n                ch = re.sub(r'[^A-Za-z0-9]', '', last)\n                if ch:\n                    clist.append(ch)\n            chain2 = clist\n    if chain1 is None or chain2 is None:\n        raise ValueError(\"Failed to parse chain orders from USalign output\")\n    return chain1, chain2\n\n# ---------------------\n# Main scoring function (no try/except, no fallbacks)\n# ---------------------\ndef score(solution: pd.DataFrame, submission: pd.DataFrame, row_id_column_name: str, usalign_bin_hint: str = None) -> float:\n    \"\"\"\n    Enhanced scoring with chain-permutation handling for multicopy targets.\n    This version contains no try/except blocks and will raise on any error.\n    \"\"\"\n    # determine usalign binary\n    if usalign_bin_hint:\n        usalign_bin = usalign_bin_hint\n    else:\n        if os.path.exists('/kaggle/input/usalign/USalign') and not os.path.exists('/kaggle/working/USalign'):\n            shutil.copy2('/kaggle/input/usalign/USalign', '/kaggle/working/USalign')\n            os.chmod('/kaggle/working/USalign', 0o755)\n        usalign_bin = '/kaggle/working/USalign' if os.path.exists('/kaggle/working/USalign') else 'USalign'\n\n    sol = solution.copy()\n    sub = submission.copy()\n    sol['target_id'] = sol['ID'].apply(lambda x: '_'.join(str(x).split('_')[:-1]))\n    sub['target_id'] = sub['ID'].apply(lambda x: '_'.join(str(x).split('_')[:-1]))\n\n    results = []\n\n    for target_id, group_native in sol.groupby('target_id'):\n        group_predicted = sub[sub['target_id'] == target_id]\n        has_chain_copy = ('chain' in group_native.columns) and ('copy' in group_native.columns)\n        is_multicopy = has_chain_copy and (group_native['copy'].astype(float).max() > 1)\n\n        # precompute native models that have coords\n        native_with_coords = []\n        for native_cnt in range(1, 41):\n            native_pdb = f'native_{target_id}_{native_cnt}.pdb'\n            resolved_native = write2pdb(group_native, native_cnt, native_pdb)\n            if resolved_native > 0:\n                native_with_coords.append(native_cnt)\n            else:\n                if os.path.exists(native_pdb):\n                    os.remove(native_pdb)\n\n        if not native_with_coords:\n            raise ValueError(f\"No native models with coordinates for target {target_id}\")\n\n        best_per_pred = []\n        for pred_cnt in range(1, 6):\n            if not is_multicopy:\n                predicted_pdb = f'predicted_{target_id}_{pred_cnt}.pdb'\n                resolved_pred = write2pdb(group_predicted, pred_cnt, predicted_pdb)\n                if resolved_pred <= 2:\n                    #print(f\"Predicted model {pred_cnt} for target {target_id} has insufficient coordinates\")\n                    best_per_pred.append( 0.0 )\n                    continue\n                \n                scores = []\n                for native_cnt in native_with_coords:\n                    native_pdb = f'native_{target_id}_{native_cnt}.pdb'\n                    out = run_usalign_raw(predicted_pdb, native_pdb, usalign_bin=usalign_bin, align_sequence=False, tmscore=1)\n                    s = parse_tmscore_output(out)\n                    scores.append(s)\n                best_per_pred.append(max(scores))\n\n            else:\n                # multicopy\n                # strict: require chain and copy columns convertible\n                gn_sorted = group_native.copy()\n                gn_sorted['__resid_int'] = gn_sorted['resid'].astype(int)\n                gn_sorted = gn_sorted.sort_values('__resid_int').reset_index(drop=True)\n                groups_list = []\n                for _, r in gn_sorted.iterrows():\n                    chain_val = r['chain']\n                    copy_i = int(r['copy'])\n                    groups_list.append((chain_val, copy_i))\n\n                # predicted multichain - groups_list must match predicted residue count or error\n                dfp_sorted = group_predicted.copy()\n                dfp_sorted['__resid_int'] = dfp_sorted['resid'].astype(int)\n                dfp_sorted = dfp_sorted.sort_values('__resid_int').reset_index(drop=True)\n                if len(groups_list) != len(dfp_sorted):\n                    raise ValueError(f\"groups_list length ({len(groups_list)}) does not match predicted residue count ({len(dfp_sorted)}) for target {target_id}\")\n\n                predicted_multi_pdb = f'pred_multi_{target_id}_{pred_cnt}.pdb'\n                resolved_pred_multi, pred_chain_letters = write2pdb_multichain_from_groups(group_predicted, pred_cnt, predicted_multi_pdb, groups_list)\n                if resolved_pred_multi == 0:\n                    #print(f\"Predicted multi model {pred_cnt} for target {target_id} has no coordinates\")\n                    best_per_pred.append( 0.0 )\n                    continue\n\n                scores = []\n                for native_cnt in native_with_coords:\n                    native_multi_pdb = f'native_multi_{target_id}_{native_cnt}.pdb'\n                    resolved_native_multi = write2pdb_multichain_from_solution(group_native, native_cnt, native_multi_pdb)\n                    if resolved_native_multi == 0:\n                        continue\n\n                    raw_out = run_usalign_raw(predicted_multi_pdb, native_multi_pdb, usalign_bin=usalign_bin, align_sequence=True, tmscore=0)\n                    chain1, chain2 = parse_usalign_chain_orders(raw_out)  # will raise if parsing fails\n\n                    # build native->pred mapping chain2[i] -> chain1[i]\n                    native_to_pred = {n_ch: p_ch for n_ch, p_ch in zip(chain2, chain1)}\n\n                    # canonical native order = chain2 unique in order seen\n                    #native_chain_order = []\n                    #for ch in chain2:\n                    #    if ch not in native_chain_order:\n                    #        native_chain_order.append(ch)\n                    native_chain_order = list(native_to_pred.keys())\n                    native_chain_order.sort() # this is critical...\n\n                    # predicted chain order by following native chain A,B,...\n                    pred_chain_order = [native_to_pred[n_ch] for n_ch in native_chain_order if native_to_pred.get(n_ch) is not None]\n\n                    # construct pred_positions_by_chain\n                    pred_positions_by_chain = {}\n                    for idx, ch in enumerate(pred_chain_letters):\n                        if ch is None:\n                            continue\n                        pred_positions_by_chain.setdefault(ch, []).append(idx)\n\n                    # require that each chain in pred_chain_order exists in pred_positions_by_chain\n                    pred_chain_order = [p for p in pred_chain_order if p in pred_positions_by_chain]\n\n                    # form permuted indices by concatenation\n                    permuted_indices = []\n                    for ch in pred_chain_order:\n                        permuted_indices.extend(pred_positions_by_chain[ch])\n                    # append any remaining\n                    for idx in range(len(pred_chain_letters)):\n                        if idx not in permuted_indices:\n                            permuted_indices.append(idx)\n\n                    # write permuted single-chain predicted and native single-chain\n                    pred_single_perm = f'pred_permuted_{target_id}_{pred_cnt}_{native_cnt}.pdb'\n                    written_pred_single = write2pdb_singlechain_permuted_pred(group_predicted, pred_cnt, permuted_indices, pred_single_perm)\n                    native_single = f'native_single_{target_id}_{native_cnt}.pdb'\n                    written_native = write2pdb_singlechain_native(group_native, native_cnt, native_single)\n\n                    if written_pred_single <= 2 or written_native <= 2:\n                        raise ValueError(f\"Insufficient residues after permutation for target {target_id}, pred {pred_cnt}, native {native_cnt}\")\n\n                    out = run_usalign_raw(pred_single_perm, native_single, usalign_bin=usalign_bin, align_sequence=False, tmscore=1)\n                    score_final = parse_tmscore_output(out)\n                    scores.append(score_final)\n\n                best_per_pred.append(max(scores))\n\n        results.append(max(best_per_pred))\n\n    if not results:\n        pass\n        #raise ValueError(\"No targets scored\")\n    return float(sum(results) / len(results)) if len(results)>0 else 0.0","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import runpy\nmodule_globals = runpy.run_path(\"/kaggle/working/metric.py\")\nscore = module_globals['score']","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if not IS_SCORING_RUN:\n    import pandas as pd\n    sub = pd.read_csv('/kaggle/working/submission.csv')\n    sol = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/validation_labels.csv')\n\n    sub['target_id'] = sub['ID'].apply(lambda x: '_'.join(str(x).split('_')[:-1]))\n    sol['target_id'] = sol['ID'].apply(lambda x: '_'.join(str(x).split('_')[:-1]))\n    \n    # Get unique targets from submission\n    sub_targets = sub['target_id'].unique()\n    \n    results = []\n    for target_id in sub_targets:\n        group_native = sol[sol['target_id'] == target_id]\n        group_predicted = sub[sub['target_id'] == target_id]\n        result = score(group_native, group_predicted, 'ID')\n        print(f\"{target_id}: {result:.4f}\")\n        results.append(result)\n    \n    print(f\"\\nMean score: {sum(results)/len(results):.4f} (n={len(results)})\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}