{"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":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":16320058,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":14604295,"datasetId":9328538,"databundleVersionId":15440074},{"sourceType":"datasetVersion","sourceId":14962460,"datasetId":9577079,"databundleVersionId":15833819},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"datasetVersion","sourceId":14962495,"datasetId":9577097,"databundleVersionId":15833858},{"sourceType":"datasetVersion","sourceId":15480097,"datasetId":9611068,"databundleVersionId":16404042},{"sourceType":"datasetVersion","sourceId":10855324,"datasetId":6742586,"databundleVersionId":11219268}],"dockerImageVersionId":31287,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"USE_PROTENIX = True","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Setup","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nfrom pathlib import Path\n\n!pip install /kaggle/input/datasets/gabalz/rnastruct/parasail-1.3.4-py2.py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl --no-deps\n!pip install /kaggle/input/datasets/gabalz/rnastruct/viennarna-2.7.2-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl --no-deps\n\nPROTENIX_OUT_DIR = Path(\"/kaggle/working/protenix\")\nif USE_PROTENIX:\n    # Protenix setup is from: https://www.kaggle.com/code/sigmaborov/stanford-rna-3d-folding-top-1-solution\n    !pip install --no-index --no-deps /kaggle/input/datasets/kami1976/biopython-cp312/biopython-1.86-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl\n    !pip install --no-index --no-deps /kaggle/input/datasets/amirrezaaleyasin/biotite/biotite-1.6.0-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl\n    !pip install --no-index --no-deps /kaggle/input/datasets/amirrezaaleyasin/rdkit-2025-9-5/rdkit-2025.9.5-cp312-cp312-manylinux_2_28_x86_64.whl\n    os.environ[\"LAYERNORM_TYPE\"] = \"torch\"\n    # os.environ.setdefault(\"RNA_MSA_DEPTH_LIMIT\", \"512\")\n    PROTENIX_DIR = Path(\n        \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted\"\n        \"/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1\"\n    )\n    PROTENIX_OUT_DIR.mkdir(parents=True, exist_ok=True)\n    os.environ[\"PROTENIX_ROOT_DIR\"] = str(PROTENIX_DIR)\n    sys.path.append(str(PROTENIX_DIR))\n\nPROTENIX_MODEL = \"protenix_base_20250630_v1.0.0\"\nPROTENIX_MAX_LEN = 768\nPROTENIX_DTYPE = 'fp32'\n\nPROTENIX_USE_MSAS = [True, True]\nPROTENIX_USE_RNA_MSAS = [True, True]\nPROTENIX_USE_TEMPLATES = [True, True]\nPROTENIX_N_CYCLE = 10\nPROTENIX_N_STEP = 200\n\nPROTENIX_USE_UNPAIRED_MSA_FILES = True\nPROTENIX_USE_TRUNC_MSA_FILES = True\nPROTENIX_JSON_SPLIT_CHAINS = True\n\nPROTENIX_MULTI_ALIGN_TO_BEST_TBM = False\nPROTENIX_USE_LIGAND = False\nPROTENIX_SORT_BY_PLDDT = False\nPROTENIX_MIN_N_SAMPLE = 1\nPROTENIX_ORIENT_CHUNKS = False\nPROTENIX_CHUNK_SIZE = 256  # overlap size for chunks, use None to disable\nPROTENIX_CHUNK_RAMP = True\nPROTENIX_ORDER_MULTI_SAMPLES = False\nPROTENIX_MULTI_CHAIN_SKIP = True\nPROTENIX_AS_NEEDED_MAXITERS = False\nPROTENIX_QUERY_MAXITERS = {\n    True: 5,  # single-query\n    False: 1,  # multi-query\n}\nPROTENIX_QUERY_MAXSECONDS = {\n    True: 300,  # single-query\n    False: 300,  # multi-query\n}\nSKIP_PROTENIX_TIDS = [] #['9MME']  # this parameter is only used for validation","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\nimport random\nimport numpy as np\nimport pandas as pd\nimport numba\nimport runpy\nimport shutil\nimport psutil\nimport pickle\nimport re\nimport gc\nimport math\nimport json\nimport random\nimport torch\nfrom tqdm import tqdm\nfrom collections import namedtuple\nfrom scipy.interpolate import interp1d\nfrom scipy.optimize import linear_sum_assignment\nfrom joblib import Parallel, delayed\nfrom concurrent.futures import ProcessPoolExecutor, as_completed\n\nimport parasail\nimport RNA\n\nT_NOTEBOOK_START = time.time()\nIS_COMPETITION_RERUN = os.getenv('KAGGLE_IS_COMPETITION_RERUN')\nNWORKERS = psutil.cpu_count(logical=False)\nprint(f'NWORKERS: {NWORKERS}')\n\nSEED = 137\nos.environ[\"PYTHONHASHSEED\"] = str(SEED)\nrandom.seed(SEED)\nnp.random.seed(SEED)\n\nRNG_SEED = int(np.random.randint(1e6))\nprint(f'RNG_SEED: {RNG_SEED}')\nRNG = np.random.default_rng(RNG_SEED)\nPROTENIX_SEED = int(np.random.randint(1e6))\n\nif USE_PROTENIX:\n    from protenix.utils.seed import seed_everything\n    from protenix.data.inference.infer_dataloader import InferenceDataset\n    from runner.inference import (\n        InferenceRunner,\n        update_inference_configs,\n        update_gpu_compatible_configs,\n    )\n    print(f'PROTENIX_SEED: {PROTENIX_SEED}')\n\n# T_MAX = None\nT_AVAILABLE_SECONDS = 3600 * (8 if IS_COMPETITION_RERUN else 2.5) - 600\nT_MAX = time.time() + T_AVAILABLE_SECONDS  # max notebook time with 10 minutes safety buffer\nprint(f'T_MAX: {T_MAX}, T_AVAILABLE_SECONDS: {T_AVAILABLE_SECONDS}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_DIR = Path(\"/kaggle/input/stanford-rna-3d-folding-2\")\nMSA_DIR = DATA_DIR / 'MSA'\n\nDO_VALIDATE = (not IS_COMPETITION_RERUN)\nUSE_TEMPORAL_CUTOFF = DO_VALIDATE\n\nNUM_PREDS = 5  # number of predictions\nCOORDS_MIN = -999.999\nCOORDS_MAX = 9999.999\n\nSKIPPED_VALID_TARGET_IDS = {}\n# SKIPPED_VALID_TARGET_IDS = {'9MME'}\nSS_ONLY = True\n\n# Alignment score parameters:\nAlignConfig = namedtuple('AlignConfig', [\n    'id', 'pmatch', 'pmismatch', 'smatch', 'smismatch', 'gap_open', 'gap_extend',\n])\nALIGN_CONFIGS = [\n    AlignConfig(0, 30, -20, 10, -10, 95, 4),\n    AlignConfig(1, 30, -20, 1, -1, 60, 10),\n]\n\nMAX_SEQ_LEN_MISMATCH = 1.0  # maximum sequence length mismatch for some alignments\nMAX_CHAIN_LEN_MISMATCH = 1.0  # maximum chain sequence length mismatch\nMIN_SEQMR = 0.0  # minimum global sequence match ratio\nMIN_SCORE = 0.0  # minimum global alignment ratio score\nMAX_COORDS_FILL_RATE = 0.90  # maximum fill rate for coordinates\nMIN_COORDS_DIST = 0.01  # coords will be discarded within this similarity score\nCOORDS_DELAY_TOL = None  # trying to select candidates which are further away from others than this threshold\nUSE_PROTENIX_SCORES = (0.5, 0.8)  # if the last candidates score is below these thresholds, replace them with protenix guesses\nSTEM_OC_MISMATCH_MULTIPLIER = 2\nGROUP_DIST_LIMIT = 2.0  # clustering radius of average candidate coordinates\nMSA_SEQ_COVER_RATIO = 0.1  # keep only truncated MSA sequences with higher nogaps/len than this ratio\nCAND_MULTIPLIER = 10  # multiplies the number of requested candidates, so controls the number of candidates selected for evaluation\nPC_SCORE_WEIGHT = 1.0  # 1.0: only uses the fill-adjusted score, 0.0: only uses alignment score\nUSE_INTER_CHAIN_TBM = False\nUSE_COORD_POSTPROCESSING = False\n\n# Secondary RNA structure processing parameters:\nMSA_CHUNK_SIZE = 1024\nMSA_MIN_SEQ_LEN = 10  # minimum (non-gap) residues in MSA sequences\nMSA_APPROX_THRESHOLD = 300  # use approximate Neff calculation and subsample MSA sequences if there are more than this amount of sequences\nMSA_SEQ_IDENTITY_THRESHOLD = 0.8  # MSA sequences matching beyond this limit are dropped\nMSA_NEFF_THRESHOLD = 15  # threshold to choose between single-sequence and consensus secondary structures","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data","metadata":{}},{"cell_type":"code","source":"label_dtypes = {'chain': 'str'}\n\ntrain_seqs = pd.read_csv(DATA_DIR / 'train_sequences.csv')\ntrain_labels = pd.read_csv(DATA_DIR / 'train_labels.csv', dtype=label_dtypes)\ntest_seqs = pd.read_csv(DATA_DIR / 'test_sequences.csv')\n\nvalid_seqs_path = DATA_DIR / 'validation_sequences.csv'\nvalid_labels_path = DATA_DIR / 'validation_labels.csv'\nvalid_seqs = pd.read_csv(valid_seqs_path) if os.path.exists(valid_seqs_path) else pd.DataFrame()\nvalid_labels = pd.read_csv(valid_labels_path, dtype=label_dtypes) if os.path.exists(valid_labels_path) else pd.DataFrame()\n\ndo_merge_train_and_valid_data = valid_seqs is not None and valid_labels is not None and not DO_VALIDATE\nif do_merge_train_and_valid_data:\n    print(f'Merging train and valid data, train_seqs.shape: {train_seqs.shape}.')\n    train_seqs = pd.concat([train_seqs, valid_seqs], ignore_index=True)\n    train_labels = pd.concat([train_labels, valid_labels], ignore_index=True)\ntrain_seqs['temporal_cutoff'] = train_seqs['temporal_cutoff'].fillna('1900-01-01')\nif train_seqs['target_id'].duplicated().any():\n    print(f'Warning: removing duplicated target_ids from train_seqs!')\n    train_seqs = train_seqs.drop_duplicates(subset='target_id', keep='first')\ntrain_seqs = train_seqs.set_index('target_id')\n\nprint('Number of records:\\n'\n      f'    train, seqs:{len(train_seqs)}, labels:{len(train_labels)}\\n'\n      f'    valid, seqs:{len(valid_seqs)}, labels:{len(valid_labels)}\\n'\n      f'    test,  seqs:{len(test_seqs)}')\ntrain_seqs.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def process_labels(labels_df):\n    coords_dict = {}\n    prefixes = labels_df[\"ID\"].str.rsplit(\"_\", n=1).str[0]\n    for id_prefix, group in labels_df.groupby(prefixes, sort=False):\n        coords = group.sort_values(\"resid\")[[\"x_1\", \"y_1\", \"z_1\"]].values\n        coords_nonan = np.nan_to_num(coords)\n        coords[(coords_nonan < COORDS_MIN).any(axis=1)] = np.nan\n        coords[(coords_nonan > COORDS_MAX).any(axis=1)] = np.nan\n        coords_dict[id_prefix] = coords\n    return coords_dict\n\n\ntrain_coords_dict = process_labels(train_labels)\n\nprint('\\nSequence length statistics:')\ndisplay(pd.Series([len(v) for v in train_coords_dict.values()]).describe())\nprint('\\nNaN count statistics:')\ndisplay(pd.Series([np.count_nonzero(np.isnan(v)) for v in train_coords_dict.values()]).describe())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Coordinate operations","metadata":{}},{"cell_type":"code","source":"def kabsch(P, Q):\n    \"\"\"\n    Computes the optimal rotation matrix R and translation vector t\n    that minimizes RMSD between two paired point sets P and Q as Q ~ P @ R.T + t.\n    \"\"\"\n    centroid_P = np.nanmean(P, axis=0)\n    centroid_Q = np.nanmean(Q, axis=0)\n    P_centered = P - centroid_P\n    Q_centered = Q - centroid_Q\n\n    # Remove rows containing NaN\n    mask = ~(\n        np.isnan(P_centered).any(axis=1) |\n        np.isnan(Q_centered).any(axis=1)\n    )\n    P_centered = P_centered[mask]\n    Q_centered = Q_centered[mask]\n\n    H = P_centered.T @ Q_centered\n    U, S, Vt = np.linalg.svd(H)\n    R = Vt.T @ U.T\n\n    if np.linalg.det(R) < 0:\n        Vt[-1, :] *= -1\n        R = Vt.T @ U.T\n\n    t = centroid_Q - R @ centroid_P\n    return R, t\n\n\ndef nan_to_mean(arr):\n    arr = arr.copy()\n    column_means = np.nanmean(arr, axis=0)\n    nan_mask = np.isnan(arr)\n    arr[nan_mask] = np.take(column_means, np.where(nan_mask)[1])\n    return arr\n\n\ndef calc_coords_dist(c, oc):\n    c = nan_to_mean(c)\n    oc = nan_to_mean(oc)\n    R, t = kabsch(c, oc)\n    c = c @ R.T + t\n    return np.sqrt(np.mean(np.sum((c - oc)**2, axis=1)))\n\n\ndef calc_mean_coords(cs):\n    ms = [nan_to_mean(c) for c in cs]\n    tc = [cs[0]]\n    mc0 = nan_to_mean(tc[0])\n    for c in cs[1:]:\n        mc = nan_to_mean(c)\n        R, t = kabsch(mc, mc0)\n        tc.append(c @ R.T + t)\n    return np.nanmean(np.stack(tc, axis=2), axis=2)\n\n\ndef calc_nmissing_coords(coords):\n    return np.sum(np.isnan(coords).any(axis=1))\n\n\ndef get_filled_coords(coords, target_dist=5.0, iterations=10, k_spring=0.1, tol=1e-5):\n    nan_mask = np.isnan(coords).any(axis=1)\n    nmissing = np.sum(nan_mask)\n    if nmissing == 0:\n        return coords\n    if nmissing == len(coords):\n        return None\n    is_anchor = ~nan_mask\n    \n    # Initial linear interpolation\n    valid_idx = np.where(is_anchor)[0]\n    f = interp1d(valid_idx, coords[valid_idx], axis=0, kind='linear', fill_value=\"extrapolate\")\n    p = f(np.arange(len(coords)))\n\n    for i in range(iterations):\n        p_old = p.copy()\n\n        # 1. Vectorized Distance Constraint (Forward/Backward)\n        diffs = np.diff(p, axis=0)\n        dists = np.linalg.norm(diffs, axis=1, keepdims=True)\n        dists[dists < 1e-4] = 1e-4\n        \n        # Displacement needed for each bond\n        errors = (dists - target_dist) / dists\n        corrections = diffs * 0.5 * errors\n        \n        # Apply corrections to non-anchor atoms\n        p[:-1][~is_anchor[:-1]] += corrections[~is_anchor[:-1]]\n        p[1:][~is_anchor[1:]] -= corrections[~is_anchor[1:]]\n\n        # 2. Vectorized Laplacian Smoothing\n        # p_i = p_i + k * ( (p_{i-1} + p_{i+1})/2 - p_i )\n        avg_neighbors = (p[:-2] + p[2:]) / 2.0\n        p[1:-1][~is_anchor[1:-1]] += k_spring * (avg_neighbors[~is_anchor[1:-1]] - p[1:-1][~is_anchor[1:-1]])\n\n        # Early Exit\n        if np.max(np.linalg.norm(p - p_old, axis=1)) < tol:\n            break\n            \n    return p","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# postprocessing\nimport torch.nn.functional as F\n\n\ndef parse_dot_bracket(dot_bracket):\n    stack = []\n    pairs = []\n    for i, c in enumerate(dot_bracket):\n        if c == \"(\":\n            stack.append(i)\n        elif c == \")\":\n            j = stack.pop()\n            pairs.append((j, i))\n    return pairs\n\n\ndef build_backbone_pairs(segments):\n    pairs = []\n    for begin, end, _ in segments:\n        for i in range(begin, end-1):\n            pairs.append((i, i + 1))\n    return pairs\n\n\ndef refine_rna_coords(\n    query_seq,\n    query_segments,\n    query_db,\n    coords,\n    steps=30,\n    lr=0.03,\n    device=None,\n    verbose=False,\n):\n    if device is None:\n        device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    coords_pred = torch.tensor(coords, dtype=torch.float32, device=device)\n    coords_ref = coords_pred.clone().detach().requires_grad_(True)\n\n    N = len(query_seq)\n    # assert coords_ref.shape == (N, 3), f'{coords_ref.shape} != ({N}, 3), coords.shape: {coords.shape}'\n\n    backbone_pairs = build_backbone_pairs(query_segments)\n    base_pairs = parse_dot_bracket(query_db)\n\n    backbone_pairs = torch.tensor(backbone_pairs, device=device)\n    base_pairs = torch.tensor(base_pairs, device=device) if base_pairs else None\n\n    optimizer = torch.optim.Adam([coords_ref], lr=lr)\n    for step in range(steps):\n        optimizer.zero_grad()\n\n        # -------------------------\n        # Backbone constraint (soft range)\n        # ideal ≈ 6 Å\n        # allow tolerance\n        # -------------------------\n        if len(backbone_pairs) > 0:\n            p1 = coords_ref[backbone_pairs[:, 0]]\n            p2 = coords_ref[backbone_pairs[:, 1]]\n            d = torch.norm(p1 - p2, dim=1)\n\n            lower = 4.0\n            upper = 8.0\n            L_backbone = (\n                torch.relu(lower - d) ** 2 +\n                torch.relu(d - upper) ** 2\n            ).mean()\n        else:\n            L_backbone = 0\n\n        # -------------------------\n        # Base pair constraint\n        # ideal ≈ 10.5 Å\n        # tolerance window\n        # -------------------------\n        if base_pairs is not None and len(base_pairs) > 0:\n            p1 = coords_ref[base_pairs[:, 0]]\n            p2 = coords_ref[base_pairs[:, 1]]\n            d = torch.norm(p1 - p2, dim=1)\n\n            lower = 8.0\n            upper = 13.0\n            L_pair = (\n                torch.relu(lower - d) ** 2 +\n                torch.relu(d - upper) ** 2\n            ).mean()\n        else:\n            L_pair = 0\n\n        # -------------------------\n        # Vectorized steric clashes\n        # -------------------------\n        dist_matrix = torch.cdist(coords_ref, coords_ref)\n        clash_mask = torch.triu(\n            torch.ones(N, N, device=device), diagonal=3\n        ).bool()\n\n        clash_dist = dist_matrix[clash_mask]\n        L_clash = torch.relu(3.0 - clash_dist) ** 2\n        L_clash = L_clash.mean() if len(L_clash) > 0 else 0\n\n        # -------------------------\n        # Backbone smoothness\n        # -------------------------\n        # L_smooth = 0.0\n        # for (b, e, _) in query_segments:\n        #     c = coords_ref[b:e, :]\n        #     v1 = torch.norm(c[1:-1] - c[:-2], dim=1)\n        #     v2 = torch.norm(c[2:] - c[1:-1], dim=1)\n        #     L_smooth += torch.mean((v2 - v1) ** 2)\n        # L_smooth /= len(query_segments)\n\n        # -------------------------\n        # Minimal movement penalty\n        # strong regularizer\n        # -------------------------\n        L_dev = torch.mean((coords_ref - coords_pred) ** 2)\n\n        # -------------------------\n        # Total loss\n        # -------------------------\n        loss = (\n            1.0 * L_backbone +\n            0.1 * L_pair +\n            # 0.1 * L_smooth +\n            4.0 * L_clash +\n           10.0 * L_dev \n        )\n        loss.backward()\n        optimizer.step()\n\n        if verbose and step % 50 == 0:\n            print(f\"step {step} loss {loss.item():.4f}\")\n\n    return coords_ref.detach().cpu().numpy()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## RNA structure","metadata":{}},{"cell_type":"code","source":"def parse_fasta(fasta_content):\n    \"\"\"\n    Parse FASTA content into dictionary.\n\n    Args:\n        fasta_content: Multi-line FASTA string with format:\n        >1A1T_1|Chain A[auth B]|SL3 STEM-LOOP RNA|\n        or\n        >104D_1|Chains A[auth A], B[auth B]|DNA/RNA (...)|\n\n    Returns:\n        Dictionary mapping auth chain_id to (sequence, list_of_auth_chain_ids)\n        Example: {\"A\": (\"ACGT\", [\"A\", \"B\"]), \"C\": (\"UGCA\", [\"C\"])}\n        The key is the auth chain ID, and the list contains all auth chain IDs for this sequence\n    \"\"\"\n    result = {}\n    lines = fasta_content.strip().split(\"\\n\")\n\n    i = 0\n    while i < len(lines):\n        line = lines[i].strip()\n\n        if line.startswith(\">\"):\n            # Parse new format header: >104D_1|Chains A[auth A], B[auth B]|...| or >1A1T_1|Chain A[auth B]|...|\n            # Extract the chains part (between first | and second |)\n            parts = line.split(\"|\")\n            if len(parts) < 2:\n                print(\"Warning: Malformed FASTA header:\", line)\n                auth_chain_ids = []\n                chains_part = \"\"\n            else:\n                chains_part = parts[1].strip()\n\n                # Extract auth chain IDs from patterns like \"Chain A[auth B]\" or \"Chains A[auth A], B[auth B] or just \"Chain A\" or \"Chains A, B\"\n                auth_chain_ids = []\n                replaced_chains_part = re.sub(r\"^Chains? \", \"\", chains_part)\n                chains = replaced_chains_part.split(\",\")\n                for chain in chains:\n                    auth_match = re.search(r\"\\[auth ([^\\]]+)\\]\", chain)\n                    if auth_match:\n                        auth_chain_ids.append(auth_match.group(1).strip())\n                    else:\n                        c = chain.strip()\n                        if c:\n                            auth_chain_ids.append(c)\n\n            if not auth_chain_ids:\n                print(\"Warning: Empty chains part:\", chains_part)\n                primary_auth_chain = None\n            else:\n                # Use the first auth chain ID as the key\n                primary_auth_chain = auth_chain_ids[0]\n\n            # Read sequence (next lines until next header or end)\n            sequence = \"\"\n            while (i + 1) < len(lines) and lines[i + 1].startswith(\">\") is False:\n                sequence += lines[i + 1].strip()\n                i += 1\n            result[primary_auth_chain] = (sequence, auth_chain_ids)\n\n        i += 1\n\n    return result\n\n\ndef parse_stoichiometry(stoich):\n    if pd.isna(stoich) or str(stoich).strip() == \"\":\n        return []\n    out = []\n    for part in str(stoich).split(';'):\n        ch, cnt = part.split(':')\n        out.append((ch.strip(), int(cnt)))\n    return out\n\n\ndef parse_chain_segments(row):\n    eq = row['sequence']\n    stoich = row.get('stoichiometry', '')\n    all_seq = row.get('all_sequences', '')\n    if pd.isna(stoich) or pd.isna(all_seq) or str(stoich).strip() == \"\" or str(all_seq).strip() == \"\":\n        return {}, []\n    try:\n        chain_dict = {k: v[0] for k, v in parse_fasta(all_seq).items()}\n        order = parse_stoichiometry(stoich)\n        return chain_dict, order\n    except Exception as e:\n        print(f'Error in parse_chain_segments, tid: {row[\"target_id\"]}, err: {e}!')\n        sys.stdout.flush()\n    return {}, []\n\n\ndef get_chain_segments(row, include_ids=True):\n    seq_len = len(row['sequence'])\n    chain_dict, order = parse_chain_segments(row)\n    single_chain_result = [(0, seq_len, '') if include_ids else (0, seq_len)]\n    if len(chain_dict) == 0 or len(order) == 0:\n        return single_chain_result\n    pos = 0\n    segs = []\n    for ch, cnt in order:\n        base = chain_dict.get(ch)\n        if base is None:\n            return single_chain_result\n        for _ in range(cnt):\n            L = len(base)\n            segs.append((pos, pos + L, ch) if include_ids else (pos, pos + L))\n            pos += L\n    if pos != seq_len:\n        return single_chain_result\n    return segs","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"rnastruct_train = pd.read_csv('/kaggle/input/datasets/gabalz/rnastruct/rnastruct_train_full.csv', index_col=0)\nif do_merge_train_and_valid_data:\n    rnastruct_valid = pd.read_csv('/kaggle/input/datasets/gabalz/rnastruct/rnastruct_valid_full.csv', index_col=0)\n    rnastruct_train = pd.concat([rnastruct_train, rnastruct_valid], ignore_index=True)\n\nif rnastruct_train['target_id'].duplicated().any():\n    print(f'Warning: removing duplicated target_ids from rnastruct_train!')\n    rnastruct_train = rnastruct_train.drop_duplicates(subset='target_id', keep='first')\nrnastruct_train = rnastruct_train.set_index('target_id')\nfor col in ['segments', 'ss_dotbrackets', 'neffs', 'cs_dotbrackets']:\n    rnastruct_train[col] = rnastruct_train[col].apply(eval)\nprint(f'rnastruct_train.shape: {rnastruct_train.shape}')\nrnastruct_train.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@numba.njit\ndef sequence_identity(seq1, seq2):\n    \"\"\"\n    Compute pairwise identity ignoring positions\n    where both sequences have gaps.\n    \"\"\"\n    matches = 0\n    valid = 0\n    for a, b in zip(seq1, seq2):\n        if a == '-' and b == '-':\n            continue\n        if a != '-' and b != '-':\n            valid += 1\n            if a == b:\n                matches += 1\n    if valid == 0:\n        return 0.0\n    return matches / valid\n\n\n@numba.njit\ndef compute_neff(alignment, identity_threshold=MSA_SEQ_IDENTITY_THRESHOLD):\n    \"\"\"Exact Neff (effective sequence number).\"\"\"\n    n = len(alignment)\n    weights = np.zeros(n)\n    for i in range(n):\n        count = 0\n        for j in range(n):\n            id_ij = sequence_identity(alignment[i], alignment[j])\n            if id_ij >= identity_threshold:\n                count += 1\n        weights[i] = 1.0 / count if count > 0 else 0.0\n    neff = np.sum(weights)\n    return int(np.round(neff))\n\n\n@numba.njit\ndef compute_neff_greedy(alignment, identity_threshold=MSA_SEQ_IDENTITY_THRESHOLD):\n    \"\"\"Approximate Neff (effective sequence number) using greedy clustering.\"\"\"\n    representatives = []\n    for seq in alignment:\n        for rep in representatives:\n            if sequence_identity(seq, rep) >= identity_threshold:\n                break\n        else:\n            representatives.append(seq)\n    return len(representatives), representatives\n\n\ndef calc_dotbrackets(seq, segments, chunk_len=MSA_CHUNK_SIZE):\n    dotbrackets = []\n    for (b, e, _) in segments:\n        is_single = isinstance(seq, str)\n        if is_single:\n            segment_seq = seq[b:e]\n            dotbracket = '.' * (e - b)\n        else:\n            segment_seq = [s[b:e] for s in seq]\n            dotbracket = None\n        try:\n            if len(segment_seq) > 0:\n                if is_single:\n                    dotbracket, mfe = RNA.fold_compound(segment_seq).mfe()\n                else:\n                    seq_len = len(segment_seq[0])\n                    if seq_len > chunk_len:\n                        res = []\n                        prev_i = 0\n                        nchunks = int(np.ceil(seq_len / chunk_len))\n                        for i in np.linspace(0, seq_len, 1+nchunks, dtype=int)[1:]:\n                            seq_i = [seq[prev_i:i] for seq in segment_seq]\n                            dotbracket, mfe = RNA.fold_compound(seq_i).mfe()\n                            res.append(dotbracket)\n                            prev_i = i\n                        dotbracket = ''.join(res)\n                    else:\n                        dotbracket, mfe = RNA.fold_compound(segment_seq).mfe()\n        except Exception as e:\n            print(f'RNA error, target_id: {target_id}, err: {e}')\n            sys.stdout.flush()\n        dotbrackets.append(dotbracket)\n    return dotbrackets\n\n\n@numba.njit\ndef _msa_sequence_distance(seq1, seq2):\n    return 1.0 - sequence_identity(seq1, seq2)\n\n\n@numba.njit\ndef _msa_farthest_first_selection(sequences, k):\n    \"\"\"\n    Select k maximally diverse sequences\n    using farthest-first traversal.\n    \"\"\"\n    if len(sequences) <= k:\n        return sequences\n    selected = [sequences[0]]\n    min_dists = np.array([0.0] + [_msa_sequence_distance(seq, selected[0])\n                                  for seq in sequences[1:]])\n    while len(selected) < k:\n        maxi = np.argmax(min_dists)\n        next_seq = sequences[maxi]\n        selected.append(next_seq)\n        for i, seq in enumerate(sequences):\n            md = min_dists[i]\n            if md == 0.0:\n                continue\n            d = 0.0 if seq == next_seq else _msa_sequence_distance(seq, next_seq)\n            if d < md:\n                min_dists[i] = d\n    return selected\n\n\ndef _filter_msa_segment_seqs(segment_seqs, approx_threshold, msa_seq_cover_ratio):\n    if len(segment_seqs) <= approx_threshold:\n        return segment_seqs\n    query = segment_seqs[0]\n    segment_seqs = segment_seqs[1:]\n    lst = sorted([(len(seq.replace('-', ''))/len(seq), seq) for seq in segment_seqs],\n                 key=lambda v: v[0], reverse=True)\n    if lst[approx_threshold-1][0] <= msa_seq_cover_ratio:\n        segment_seqs = [query] + [x[1] for x in lst[:approx_threshold-1]]\n    return _msa_farthest_first_selection(segment_seqs, approx_threshold)\n\n\ndef build_rnastruct_worker(idx, row, save_dir=None, msa_dir=MSA_DIR,\n                           min_seq_len=MSA_MIN_SEQ_LEN,\n                           seq_cover_ratio=MSA_SEQ_COVER_RATIO,\n                           approx_threshold=MSA_APPROX_THRESHOLD):\n    t0 = time.time()\n    target_id = row[\"target_id\"]\n    seq = row['sequence']\n    segments = get_chain_segments(row)\n    ss_dotbrackets = calc_dotbrackets(seq, segments)\n\n    sequences = read_msa_content(msa_dir / f'{target_id}.MSA.fasta',\n                                 sequences_only=True)\n    if len(sequences) > 0 and sequences[0] != seq:\n        sequences = [seq] + sequences\n\n    if SS_ONLY:\n        neffs = [0] * len(segments)\n        neff_types = 'E' * len(segments)\n        cs_dotbrackets = ['.' * (e-b) for (b, e, _) in segments]\n    else:\n        neffs = []\n        neff_types = []\n        cs_dotbrackets = []\n        for (b, e, _) in segments:\n            segment_seqs = [seq[b:e] for seq in sequences]\n            segment_seqs = [seq for seq in segment_seqs\n                            if len(seq.replace('-', '')) > min_seq_len]\n            is_approx = len(segment_seqs) > approx_threshold\n            if len(segment_seqs) == 0:\n                neff = 0\n            else:\n                neff, representatives = (\n                    compute_neff_greedy(segment_seqs)\n                    if is_approx else (compute_neff(segment_seqs), []))\n            neffs.append(neff)\n            neff_types.append('A' if is_approx else 'E')\n    \n            if is_approx:\n                segment_seqs = _filter_msa_segment_seqs(\n                    representatives, approx_threshold, seq_cover_ratio,\n                )\n            if len(segment_seqs) == 0:\n                cs_dotbrackets.append(None)\n            else:\n                cs_dotbrackets.append(calc_dotbrackets(segment_seqs, [(0, e-b, '')])[0])\n        neff_types = ''.join(neff_types)\n\n    df = pd.DataFrame([(target_id, segments, ss_dotbrackets,\n                        neffs, neff_types, cs_dotbrackets)],\n                      index=[idx], columns=['target_id', 'segments', 'ss_dotbrackets',\n                                            'neffs', 'neff_types', 'cs_dotbrackets'])\n    if save_dir is not None:\n        df.to_csv(f'{save_dir}/{target_id}.dotbracket.csv')\n    return df\n\n\ndef build_rnastruct(data, nworkers=NWORKERS, save_dir=None):\n    rnastruct = list(tqdm(Parallel(n_jobs=nworkers, return_as='generator')(\n        delayed(build_rnastruct_worker)(idx, row.copy(), save_dir=save_dir)\n        for idx, row in data.iterrows()\n    ), desc=\"Building rnastruct\", total=len(data), position=0))\n    rnastruct = pd.concat(rnastruct).sort_index()\n    return rnastruct","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# MSA file utilities\n\n@numba.njit\ndef _process_msa_content(msa_content):\n    descriptions = []\n    sequences = []\n    for line in msa_content:\n        line = line.strip()\n        if not line or line.startswith('#'):\n            continue\n        if line.startswith('>'):\n            descriptions.append(line)\n            sequences.append(\"\")\n        else:\n            sequences[-1] += line\n    return descriptions, sequences\n\n\ndef read_msa_content(msa_file, sequences_only=False):\n    msa_content = []\n    with open(msa_file, 'rt') as f:\n        msa_content = f.readlines()\n    if len(msa_content) == 0:\n        return []\n    descs, seqs = _process_msa_content(msa_content)\n    seq_len = len(seqs[0])\n    msa_content = [(d, s) for d, s in zip(descs, seqs) if len(s) == seq_len]\n    if sequences_only:\n        return [r[1] for r in msa_content]\n    return msa_content\n\n\ndef _process_msa_output(descriptions, sequences, output_msa_file):\n    try:\n        output = [(d, s) for d, s in zip(descriptions, sequences)]\n        if len(output) <= 1:\n            return False, output\n    \n        if output_msa_file is not None:\n            with open(output_msa_file, 'wt') as f:\n                for header, seq in output:\n                    f.write(header + '\\n')\n                    f.write(seq + '\\n')\n        return True, output\n    except Exception as exc:\n        print(f'ERROR! _process_msa_output failed: {exc}!')\n        sys.stdout.flush()\n    return False, []\n\n\ndef truncate_msa_file(input_msa_file, segment_begin, segment_end, chain_id,\n                      min_cover_ratio=MSA_SEQ_COVER_RATIO, output_msa_file=None):\n    descriptions = []\n    sequences = []\n    try:\n        chain_str = f'|chain={chain_id}'\n        seen = set()\n        with open(input_msa_file, 'rt') as f:\n            header = None\n            for line in f.readlines():\n                line = line.strip()\n                if not line or line.startswith('#'):\n                    continue\n                elif line.startswith('>'):\n                    header = line\n                    if chain_str in line:\n                        header = re.sub(r'\\|copies=[0-9]+', '', line)\n                    descriptions.append(header)\n                    sequences.append(\"\")\n                elif header is not None:\n                    seq = line[segment_begin:segment_end]\n                    if (len(seq) > 0\n                            and len(seq.replace('-', ''))/len(seq) > min_cover_ratio\n                            and seq not in seen):\n                        sequences[-1] += seq\n                        seen.add(seq)\n    except Exception as exc:\n        print(f'ERROR! truncate_msa_file failed: {exc}!')\n        sys.stdout.flush()\n    return _process_msa_output(descriptions, sequences, output_msa_file)\n\n\ndef column_entropy(column):\n    counts = {}\n    for c in column:\n        counts[c] = counts.get(c, 0) + 1\n    total = len(column)\n    H = 0\n    for n in counts.values():\n        p = n/total\n        H -= p * math.log2(p)\n    return H\n\n\ndef all_chain_truncate_msa_file(input_msa_file, segments,\n                                min_cover_ratio=MSA_SEQ_COVER_RATIO,\n                                output_msa_file=None):\n    descriptions = []\n    sequences = []\n    try:\n        seen = set()\n        for (desc, seq) in read_msa_content(input_msa_file):\n            trunc_seq = ''.join([seq[b:e] for (b, e, _) in segments])\n            if (len(trunc_seq) > 0\n                    and len(trunc_seq.replace('-', ''))/len(trunc_seq) > min_cover_ratio\n                    and trunc_seq not in seen):\n                descriptions.append(desc)\n                sequences.append(trunc_seq)\n                seen.add(trunc_seq)\n    except Exception as exc:\n        print(f'ERROR! all_chain_truncate_msa_file failed: {exc}!')\n        sys.stdout.flush()\n    return _process_msa_output(descriptions, sequences, output_msa_file)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n\n@numba.njit\ndef convert_seq_to_idss_seq(seq, dotbrackets):\n    r = []\n    bopen =  {'A': 'P', 'U': 'S', 'C': 'D', 'G': 'I'}\n    bclose = {'A': 'R', 'U': 'T', 'C': 'E', 'G': 'J'}\n    for c, s in zip(seq, dotbrackets):\n        cc = c\n        if s == '(':\n            cc = bopen[c]\n        elif s == ')':\n            cc = bclose[c]\n        r.append(cc)\n    return ''.join(r)\n\n\ndef get_dotbrackets(row, neff_threshold=MSA_NEFF_THRESHOLD, ss_only=SS_ONLY):\n    if ss_only:\n        dotbrackets = row['ss_dotbrackets']\n    else:\n        dotbrackets = []\n        for ssdb, neff, csdb in zip(row['ss_dotbrackets'], row['neffs'], row['cs_dotbrackets']):\n            use_csdb = (csdb is not None and len(csdb) > 0 and neff >= neff_threshold)\n            dotbrackets.append(csdb if use_csdb else ssdb)\n    return ''.join(dotbrackets)\n\n\ntrain_idss_seq_dict = {\n    tid: (convert_seq_to_idss_seq(train_seqs.loc[tid, 'sequence'],\n                                  get_dotbrackets(rnastruct_train.loc[tid])),\n          rnastruct_train.loc[tid]['segments'])\n    for tid in train_seqs.index\n}\n\nrnastruct_test = build_rnastruct(test_seqs)\ndisplay(rnastruct_test.head())\n\ntest_idss_seq_dict = {\n    row['target_id']: (\n        convert_seq_to_idss_seq(row['sequence'],\n                                get_dotbrackets(rnastruct_test.loc[row_id])),\n        rnastruct_test.loc[row_id, 'segments'],\n    )\n    for row_id, row in test_seqs.iterrows()\n}","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Template alignment","metadata":{}},{"cell_type":"code","source":"# sa: structure-aware\n\n@numba.njit\ndef get_sa_nt_score(b1, b2, pmatch, pmismatch):\n    NT_DICT = {'A': 'A', 'P': 'A', 'R': 'A',\n               'U': 'U', 'S': 'U', 'T': 'U',\n               'C': 'C', 'D': 'C', 'E': 'C',\n               'G': 'G', 'I': 'G', 'J': 'G'}\n    nt1 = NT_DICT[b1]\n    nt2 = NT_DICT[b2]\n    score = 0\n    if nt1 == nt2:\n        score = pmatch\n    else:\n        score = pmismatch\n    return score\n\n\n@numba.njit\ndef get_sa_ss_score(b1, b2, smatch, smismatch):\n    SS_DICT = {'A': '.', 'P': '(', 'R': ')',\n               'U': '.', 'S': '(', 'T': ')',\n               'C': '.', 'D': '(', 'E': ')',\n               'G': '.', 'I': '(', 'J': ')'}\n    ss1 = SS_DICT[b1]\n    ss2 = SS_DICT[b2]\n    score = 0\n    if ss1 == ss2:\n        score += smatch\n    elif {ss1, ss2} == {'(', ')'}:\n        score += int(smismatch * STEM_OC_MISMATCH_MULTIPLIER)\n    else:\n        score += smismatch\n    return score    \n\n\n@numba.njit\ndef get_sa_score(b1, b2, pmatch, pmismatch, smatch, smismatch):\n    return (get_sa_nt_score(b1, b2, pmatch, pmismatch)\n            + get_sa_ss_score(b1, b2, smatch, smismatch))\n\n\ndef get_parasail_sa_distance_matrix(pmatch, pmismatch, smatch, smismatch):\n    alphabet = \"APRUSTCDEGIJ\"\n    matrix = parasail.matrix_create(alphabet, pmatch, pmismatch)\n    for i, b1 in enumerate(alphabet):\n        for j, b2 in enumerate(alphabet):\n            matrix[i, j] = get_sa_score(b1, b2, pmatch, pmismatch, smatch, smismatch)\n    return matrix\n\n\nSA_MATRICES = {\n    acfg.id: get_parasail_sa_distance_matrix(acfg.pmatch, acfg.pmismatch,\n                                             acfg.smatch, acfg.smismatch)\n    for acfg in ALIGN_CONFIGS\n}\npd.DataFrame(SA_MATRICES[ALIGN_CONFIGS[0].id].matrix, index=list(\"APRUSTCDEGIJX\"), columns=list(\"APRUSTCDEGIJX\"))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SimilarSeq = namedtuple('SimilarSeq', ['score', 'templ_tid', 'templ_seq', 'templ_segments', 'aln_mapping', 'best_mapping'])\n\n\ndef calc_score(ascore, pmatch, query_len, templ_len=np.inf):\n    return ascore / (pmatch * min(query_len, templ_len))\n\n\ndef calc_alignment_score(q, t, acfg=None):\n    if acfg is None:\n        acfg = ALIGN_CONFIGS[0]\n    align_result = parasail.nw_scan_32(q, t, acfg.gap_open, acfg.gap_extend, SA_MATRICES[acfg.id])\n    return calc_score(align_result.score, acfg.pmatch, len(q), len(t))\n\n\ndef extract_chains(seq, segments, encode=False):\n    chains = [seq[s:e] for (s, e, _) in segments]\n    if encode:\n        chains = [c.encode() for c in chains]\n    return chains\n\n\ndef find_alignment(query_seq, query_segments,\n                   templ_seq, templ_segments,\n                   acfg=None, max_chain_len_mismatch=MAX_CHAIN_LEN_MISMATCH):\n    if isinstance(query_seq, list):\n        assert len(query_seq) == len(query_segments)\n        qs = query_seq  # comes preprocessed \n    else:\n        qs = extract_chains(query_seq, query_segments)\n    ts = extract_chains(templ_seq, templ_segments)\n    nqs = len(qs)\n    nts = len(ts)\n\n    score_mat = np.full((nqs, nts), -np.inf)\n    for i, q in enumerate(qs):\n        for j, t in enumerate(ts):\n            if abs(len(q) - len(t)) / max(len(q), len(t)) <= max_chain_len_mismatch:\n                score_mat[i, j] = calc_alignment_score(q, t, acfg=acfg)\n    if np.all(np.isneginf(score_mat)):\n        return 0.0, [], []  # no valid alignments\n\n    min_score = np.nanmin(score_mat[np.isfinite(score_mat)])\n    cost_mat = np.where(\n        np.isfinite(score_mat),\n        -score_mat,\n        -min_score + 1e6,\n    )\n    try:\n        row_ind, col_ind = linear_sum_assignment(cost_mat)\n    except:\n        print(f'min_score: {min_score}')\n        print(cost_mat)\n        raise\n\n    best_mapping = []\n    for i in range(nqs):\n        if np.isfinite(score_mat[i, :]).any():\n            best_mapping.append((i, np.argmax(score_mat[i, :])))\n\n    score = 0.0\n    aln_mapping = []\n    for i, j in zip(row_ind, col_ind):\n        if np.isfinite(score_mat[i, j]):\n            score += score_mat[i, j]\n            aln_mapping.append((i, j))\n    return score, aln_mapping, best_mapping\n\n\ndef find_similar_sequences(query_seq, query_segments, templ_seq_dict,\n                           acfg=None, max_seq_len_mismatch=MAX_SEQ_LEN_MISMATCH):\n    qs = extract_chains(query_seq, query_segments)\n\n    similar_seqs = []\n    for templ_tid, (templ_seq, templ_segments) in templ_seq_dict.items():\n        if ((len(query_seq) - len(templ_seq))\n            / max(len(templ_seq), len(query_seq))) <= max_seq_len_mismatch:\n            try:\n                score, aln_mapping, best_mapping = find_alignment(\n                    qs, query_segments, templ_seq, templ_segments, acfg=acfg,\n                )\n                if len(aln_mapping) > 0:\n                    similar_seqs.append(SimilarSeq(score, templ_tid, templ_seq, templ_segments,\n                                                   aln_mapping, best_mapping))\n            except Exception as e:\n                print(f'Error in find_alignment: {e}!')\n                sys.stdout.flush()\n\n    similar_seqs.sort(key=lambda x: x.score, reverse=True)\n    return similar_seqs","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@numba.njit\ndef _init_query_coords(aligned_query, aligned_templ, query_len, templ_coords):\n    NT_DICT = {'A': 'A', 'P': 'A', 'R': 'A',\n               'U': 'U', 'S': 'U', 'T': 'U',\n               'C': 'C', 'D': 'C', 'E': 'C',\n               'G': 'G', 'I': 'G', 'J': 'G',\n               'N': 'N'}\n    coords = np.zeros((query_len, 3), dtype=float)\n    coords.fill(np.nan)\n\n    # Leave unmatched chain segments NaNs.\n    aligned_templ = aligned_templ.replace('N', '-')\n\n    # Map template coordinates to query\n    query_idx = 0\n    templ_idx = 0\n    nmatched = 0\n    naligned = 0\n    nnogap = 0\n    for i in range(len(aligned_query)):\n        query_char = aligned_query[i]\n        templ_char = aligned_templ[i]\n\n        if query_char != \"-\" and templ_char != \"-\":\n            if templ_idx < len(templ_coords):\n                coords[query_idx] = templ_coords[templ_idx]\n            templ_idx += 1\n            query_idx += 1\n            nnogap += 1\n            if NT_DICT[query_char] == NT_DICT[templ_char]:\n                nmatched += 1\n        elif query_char != \"-\" and templ_char == \"-\":\n            query_idx += 1\n        elif query_char == \"-\" and templ_char != \"-\":\n            templ_idx += 1\n        naligned += 1\n    return nmatched, nnogap, naligned, coords\n\n\ndef _get_query_coords(aligned_query, aligned_templ, query_len, templ_coords):\n    nmatched, nnogap, naligned, coords = _init_query_coords(\n        aligned_query, aligned_templ, query_len, templ_coords)\n    nmissing = np.sum(np.isnan(coords).any(axis=1))\n    return nmatched, nnogap, naligned, nmissing, coords  # get_filled_coords(coords)\n\n\n@numba.njit\ndef calc_ascore(aligned_query_list, aligned_templ_list, acfg):\n    _, pmatch, pmismatch, smatch, smismatch, gap_open, gap_extend = acfg\n\n    score = 0\n    for aligned_query, aligned_templ in zip(aligned_query_list, aligned_templ_list):\n        assert len(aligned_query) == len(aligned_templ)\n    \n        # Do not apply explicit penalty for unmatched chains, because it can invalidate all alignments.\n        # Their scores will be downscaled by the miss rate, which only lowers their priority (see 9ZCC).\n        aligned_query = aligned_query.replace('N', '-')\n        aligned_templ = aligned_templ.replace('N', '-')\n    \n        start_idx = 0\n        end_idx = len(aligned_query) - 1\n        while start_idx <= end_idx and (\n                aligned_query[start_idx] == '-'\n                or aligned_templ[start_idx] == '-'):\n            start_idx += 1        \n        while end_idx >= start_idx and (\n                aligned_query[end_idx] == '-'\n                or aligned_templ[end_idx] == '-'):\n            end_idx -= 1\n    \n        # # Penalty for unmatched chain segments (too strong):\n        # aligned_query = aligned_query.replace('N', '-')\n        # aligned_templ = aligned_templ.replace('N', '-')\n    \n        is_gap_opened = False\n        for idx in range(start_idx, end_idx+1):\n            qch = aligned_query[idx]\n            tch = aligned_templ[idx]\n            if qch == '-' or tch == '-':\n                score -= gap_extend if is_gap_opened else gap_open\n                is_gap_opened = True\n            else:\n                score += get_sa_score(qch, tch, pmatch, pmismatch, smatch, smismatch)\n                is_gap_opened = False\n    return score\n\n\nAdaptedCoords = namedtuple('AdaptedCoords', [\n    'aligned_query', 'aligned_templ',\n    'nmatched', 'nnogap', 'naligned', 'nmissing', 'coords'])\n\n\ndef adapt_template_to_query(query_seq, query_segments, similar_seq, templ_coords, acfg=None):\n    if acfg is None:\n        acfg = ALIGN_CONFIGS[0]\n    aln_mapping = dict(similar_seq.aln_mapping)\n    coords = np.zeros((len(query_seq), 3)) * np.nan\n    aligned_query_list = []\n    aligned_templ_list = []\n    nmatched = nnogap = naligned = nmissing = 0\n    for i, (qb, qe, _) in enumerate(query_segments):\n        segment_len = qe - qb\n        j = aln_mapping.get(i)\n        if j is None:\n            aligned_query_list.append('N' * segment_len)\n            aligned_templ_list.append('N' * segment_len)\n            nmissing += segment_len\n        else:\n            tb, te, _ = similar_seq.templ_segments[j]\n            q = query_seq[qb:qe].encode()\n            t = similar_seq.templ_seq[tb:te].encode()\n            result = parasail.nw_trace_scan_32(q, t, acfg.gap_open, acfg.gap_extend, SA_MATRICES[acfg.id])\n            restrace = result.traceback\n            aligned_query = restrace.query\n            aligned_templ = restrace.ref\n            aligned_query_list.append(aligned_query)\n            aligned_templ_list.append(aligned_templ)\n            _nmatched, _nnogap, _naligned, _nmissing, _coords = _get_query_coords(\n                aligned_query=aligned_query,\n                aligned_templ=aligned_templ,\n                query_len=segment_len,\n                templ_coords=templ_coords[tb:te, :],\n            )\n            nmatched += _nmatched\n            nnogap += _nnogap\n            naligned += _naligned\n            nmissing += _nmissing\n            coords[qb:qe, :] = _coords\n\n    return AdaptedCoords(\n        aligned_query_list,\n        aligned_templ_list,\n        nmatched, nnogap, naligned, nmissing,\n        coords,\n    )","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prediction","metadata":{}},{"cell_type":"code","source":"PredCoords = namedtuple('PredCoords', ['templ_tid', 'filled_score', 'seqmr', 'score', 'coords'])\n\n\ndef get_pc_score(pc):\n    return (PC_SCORE_WEIGHT * pc.filled_score + (1.0-PC_SCORE_WEIGHT) * pc.score, pc.seqmr)\n    # return (pc.filled_score, pc.seqmr, pc.score)\n\n\ndef pc_coords_dist(pc, opc):\n    return calc_coords_dist(pc.coords, opc.coords)\n\n\ndef cluster_pcs(pcs, ncenters, verbose=False, init_pcs=[]):\n    \"\"\"Returns the centers of farthest-first clustering.\"\"\"\n    if len(pcs) == 0:\n        return [], []\n    centers = init_pcs if len(init_pcs) > 0 else [pcs[0]]\n    center_dists = [np.inf] * len(centers)\n    min_dists = np.ones(len(pcs)) * np.inf\n    for c in centers:\n        for i, pc in enumerate(pcs):\n            min_dists[i] = min(min_dists[i], pc_coords_dist(c, pc))\n    for k in range(1 if len(init_pcs) == 0 else 0, min(ncenters, len(pcs))):\n        maxi = np.argmax(min_dists)\n        center_dists.append(min_dists[maxi])\n        if verbose:\n            print(f'...dist({maxi}): {min_dists[maxi]:.3f}')  #, {[f\"{d:.1f}\" for d in min_dists]}')\n        mpc = pcs[maxi]\n        centers.append(mpc)\n        maxi_dists = np.array([pc_coords_dist(mpc, pc) for pc in pcs])\n        min_dists = np.minimum(min_dists, maxi_dists)\n    return centers, center_dists\n\n\ndef aggregate_pcs(pcs, group_dist_limit=GROUP_DIST_LIMIT):\n    if len(pcs) == 0:\n        return pcs\n    groups = [(pcs[0].score, pcs[0].coords, [pcs[0]])]\n    for pc in pcs[1:]:\n        min_dist = np.inf\n        min_group_i = -1\n        for group_i, group in enumerate(groups):\n            group_score, group_coords, group_pcs = group\n            dist = calc_coords_dist(group_coords, pc.coords)\n            if dist < min_dist:\n                min_dist = dist\n                min_group_i = group_i\n\n        if min_dist > group_dist_limit:\n            groups.append((pc.score, pc.coords, [pc]))\n        else:\n            group_pcs = groups[min_group_i][2]\n            group_pcs.append(pc)\n            groups[min_group_i] = (\n                np.mean([pc.score for pc in group_pcs]),\n                calc_mean_coords([pc.coords for pc in group_pcs]),\n                group_pcs,\n            )\n    new_pcs = []\n    for group in groups:\n        group_score, group_coords, group_pcs = group\n        nmissing = np.sum(np.isnan(group_coords).any(axis=1))\n        coords_fill_rate = nmissing / len(group_coords)\n        filled_score = group_score * (1.0 - coords_fill_rate)\n        new_pcs.append(PredCoords(\n            ','.join([pc.templ_tid for pc in group_pcs]),\n            filled_score,\n            np.mean([pc.seqmr for pc in group_pcs]),\n            group_score,\n            group_coords,\n        ))\n    return sorted(new_pcs, key=get_pc_score, reverse=True)\n\n\ndef get_filled_and_missing_chains(coords, segments):\n    filled_chains = {}\n    missing_chains = {}\n    for (b, e, chain_id) in segments:\n        seg_coords = coords[b:e, :]\n        d = (\n            missing_chains if len(seg_coords) == calc_nmissing_coords(seg_coords)\n            else filled_chains\n        )\n        d.setdefault(chain_id, []).append((b, e))\n    return filled_chains, missing_chains\n\n\ndef fill_pcs_coords(pcs, segments):\n    new_pcs = []\n    for i, pc in enumerate(pcs):\n        try:\n            filled_chains, missing_chains = get_filled_and_missing_chains(pc.coords, segments)\n            if i < 3 and len(missing_chains) > 0:\n                print(f'Missing chains {i}.{pc.templ_tid},'\n                      f' missing:{\",\".join(sorted(missing_chains.keys()))},'\n                      f' filled:{\",\".join(sorted(filled_chains.keys()))}')\n            for chain_id, seg_list in missing_chains.items():\n                for k, (mb, me) in enumerate(seg_list):\n                    coords = pc.coords.copy()\n                    if chain_id in filled_chains:\n                        segs = filled_chains[chain_id]\n                        b, e = segs[k % len(segs)]\n                        ref_coords = coords[b:e, :]\n                        alpha = 0.1\n                    else:\n                        segs = [s for sl in filled_chains.values() for s in sl]\n                        b, e = segs[k % len(segs)]\n                        ref_coords = coords[b:e, :]\n                        alpha = 0.0\n                    fill_mean = np.nanmean(ref_coords, axis=0)[None, :]\n                    if alpha > 0.0:\n                        fill_mean = fill_mean + alpha * (ref_coords - fill_mean)\n                    coords[mb:me, :] = fill_mean\n                    pc = PredCoords(pc.templ_tid, pc.filled_score, pc.seqmr, pc.score, coords)\n        except Exception as exc:\n            print(f'ERROR in fill_pcs_coords: {exc}!')\n        new_pcs.append(pc)\n        \n    return [PredCoords(pc.templ_tid, pc.filled_score, pc.seqmr, pc.score,\n                       get_filled_coords(pc.coords))\n            for pc in new_pcs]\n\n\ndef print_pcs(tag, pcs, do_print_score=False):\n    data = [f'{pc.templ_tid}|{int(pc.filled_score*100)}|{int(pc.seqmr*100)}'\n            f'{f\"|{int(pc.score*100)}\" if do_print_score else \"\"}'\n            for pc in pcs]\n    print(f\"{tag} pcs%({len(pcs)}): [{' '.join(data)}]\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## TBM","metadata":{}},{"cell_type":"code","source":"def get_tbm_pred_coords(target_id, query_seq, query_segments, similar_seqs,\n                        train_coords_dict, min_score, min_seqmr, acfg=None):\n    if len(similar_seqs) == 0:\n        return []\n    if acfg is None:\n        acfg = ALIGN_CONFIGS[0]\n    similar_seq_coords = []\n    for i, similar_seq in enumerate(similar_seqs):\n        templ_tid = similar_seq.templ_tid\n        try:\n            adapted_coords = adapt_template_to_query(\n                query_seq=query_seq,\n                query_segments=query_segments,\n                similar_seq=similar_seq,\n                templ_coords=train_coords_dict[templ_tid],\n                acfg=acfg,\n            )\n            if adapted_coords.coords is None:\n                continue\n        except Exception as e:\n            print(f'Adapting coords failed, target_id: {templ_tid}, i: {i}, err: {e}'\n                  f'\\nquery_seq: {query_seq}\\nsimilar_seq: {similar_seq}!')\n            continue\n        coords_fill_rate = adapted_coords.nmissing / len(query_seq)\n        if coords_fill_rate > MAX_COORDS_FILL_RATE:\n            continue\n        assert len(adapted_coords.coords) == len(query_seq)\n        seqmr = adapted_coords.nmatched / len(query_seq)\n        # print(f'target_id: {templ_tid}, seqmr: {seqmr:.2f}, nmatched: {adapted_coords.nmatched}')\n        if seqmr < min_seqmr:\n            continue\n        ascore = calc_ascore(adapted_coords.aligned_query,\n                             adapted_coords.aligned_templ,\n                             acfg)\n        score = calc_score(ascore, acfg.pmatch, len(query_seq), len(similar_seq.templ_seq))\n        filled_score = score * (1.0 - coords_fill_rate)\n        if filled_score < min_score:\n            continue\n        similar_seq_coords.append(PredCoords(templ_tid,\n                                             filled_score,\n                                             seqmr,\n                                             score,\n                                             adapted_coords.coords))\n    return sorted(similar_seq_coords, key=get_pc_score, reverse=True)\n\n\ndef get_with_coords(data_dict):\n    new_data_dict = {}\n    for tid, data in data_dict.items():\n        if tid in train_coords_dict:\n            new_data_dict[tid] = data\n        else:\n            print(f'Missing coordinates for {tid}!')\n    return new_data_dict\n\n\ndef _align_and_fill_coords(coords, segments, new_coords, new_segments, align_coords):\n    R, t = kabsch(new_coords, np.vstack([align_coords[b:e, :] for (b, e, _) in segments]))\n    coords = coords.copy()\n    for (b, e, _1), (cb, ce, _2) in zip(segments, new_segments):\n        coords[b:e, :] = new_coords[cb:ce, :] @ R.T + t\n    return coords    \n\n\ndef get_multi_chain_tbm_predictions(\n    target_id, query_seq, query_segments,\n    npreds, templ_idss_seq_dict, align_coords,\n):\n    if len(query_segments) <= 1:\n        return []\n    chain_segment_lists = {}\n    for b, e, chain_id in query_segments:\n        chain_segment_lists.setdefault(chain_id, []).append((b, e, chain_id))\n    if len(chain_segment_lists) <= 1:\n        return []\n\n    chain_query_seq = {}\n    chain_query_segments = {}\n    chain_similar_seqs = {}\n    for chain_id, chain_segments in chain_segment_lists.items():\n        chain_seqs = []\n        chain_segs = []\n        l = 0\n        for (b, e, _) in chain_segments:\n            seq = query_seq[b:e]\n            chain_seqs.append(seq)\n            chain_segs.append((l, l+len(seq), chain_id))\n            l += len(seq)\n        seq = ''.join(chain_seqs)\n        chain_query_seq[chain_id] = seq\n        chain_query_segments[chain_id] = chain_segs\n        chain_similar_seqs[chain_id] = find_similar_sequences(\n            query_seq=seq,\n            query_segments=chain_segs,\n            templ_seq_dict=templ_idss_seq_dict,\n        )[npreds:CAND_MULTIPLIER]\n\n    chain_pcs_dict = {}\n    for chain_id, similar_seqs in chain_similar_seqs.items():\n        chain_pcs_dict[chain_id] = get_tbm_pred_coords(\n            target_id+'_'+chain_id, chain_query_seq[chain_id],\n            chain_query_segments[chain_id], similar_seqs,\n            train_coords_dict, min_score=MIN_SCORE, min_seqmr=MIN_SEQMR)\n\n    pcs = []\n    seq = ['-'] * len(query_seq)\n    for chain_id, chain_pcs in chain_pcs_dict.items():\n        chain_segments = chain_segment_lists[chain_id]\n        for (b, e, _) in chain_segments:\n            seq[b:e] = list(query_seq[b:e])\n        if len(pcs) == 0:\n            for chain_pc in chain_pcs:\n                coords = _align_and_fill_coords(\n                    np.zeros((len(query_seq), 3)) * np.nan,\n                    chain_segments,\n                    chain_pc.coords,\n                    chain_query_segments[chain_id],\n                    align_coords,\n                )\n                pcs.append(PredCoords(chain_pc.templ_tid, chain_pc.filled_score,\n                                      chain_pc.seqmr, chain_pc.score, coords))\n        else:\n            cand_pcs = []\n            for pc in pcs:\n                for chain_pc in chain_pcs:\n                    coords = _align_and_fill_coords(\n                        pc.coords,\n                        chain_segments,\n                        chain_pc.coords,\n                        chain_query_segments[chain_id],\n                        align_coords,\n                    )\n                    cand_pcs.append(PredCoords(\n                        pc.templ_tid + '_' + chain_pc.templ_tid,\n                        pc.filled_score + chain_pc.filled_score,\n                        pc.seqmr + chain_pc.seqmr,\n                        pc.score + chain_pc.score,\n                        coords,\n                    ))\n            cand_pcs = sorted(cand_pcs, key=get_pc_score, reverse=True)\n            pcs = cand_pcs[npreds:CAND_MULTIPLIER]  # beam-search\n\n    n_chains = len(chain_segment_lists)\n    return [PredCoords(pc.templ_tid, pc.filled_score / n_chains,\n                       pc.seqmr / n_chains, pc.score / n_chains, pc.coords)\n            for pc in pcs]\n\n\ndef get_tbm_predictions(row_id, npreds, temporal_cutoff=None):\n    target_id = test_seqs.loc[row_id, 'target_id']\n    query_idss_seq, query_segments = test_idss_seq_dict[target_id]\n\n    if temporal_cutoff is None:\n        templ_idss_seq_dict = train_idss_seq_dict\n    else:\n        cutoff_tids = train_seqs[train_seqs['temporal_cutoff'] < temporal_cutoff].index\n        templ_idss_seq_dict = {tid: train_idss_seq_dict[tid] for tid in cutoff_tids}\n    templ_idss_seq_dict = get_with_coords(templ_idss_seq_dict)\n\n    pcs_lists = []\n    for acfg in ALIGN_CONFIGS:\n        similar_seqs = find_similar_sequences(\n            query_seq=query_idss_seq,\n            query_segments=query_segments,\n            templ_seq_dict=templ_idss_seq_dict,\n            acfg=acfg,\n        )\n        print(f'len(similar_seqs): {len(similar_seqs)}')\n        similar_seqs = similar_seqs[:npreds*CAND_MULTIPLIER]\n        evaled_pcs = get_tbm_pred_coords(target_id, query_idss_seq, query_segments, similar_seqs,\n                                         train_coords_dict, min_score=MIN_SCORE, min_seqmr=MIN_SEQMR,\n                                         acfg=acfg)\n        print_pcs(f'evaled {acfg.id}', evaled_pcs)\n\n        if len(evaled_pcs) == 0 or not USE_INTER_CHAIN_TBM:\n            mchain_pcs = []\n        else:\n            mchain_pcs = get_multi_chain_tbm_predictions(\n                target_id, query_idss_seq, query_segments,\n                npreds, templ_idss_seq_dict,\n                align_coords=evaled_pcs[0].coords,\n            )\n            print_pcs('mchain', mchain_pcs)\n\n        pcs = aggregate_pcs(evaled_pcs + mchain_pcs)\n        pcs = fill_pcs_coords(pcs, query_segments)\n        pcs_lists.append(pcs)\n\n    pcs = []\n    inds = [0] * len(pcs_lists)\n    used_tids = set()\n    while len(pcs) < npreds:\n        pcs_len = len(pcs)\n        for i, pcs_list in enumerate(pcs_lists):\n            ind = inds[i]\n            while ind < len(pcs_list):\n                pc = pcs_list[ind]\n                ind += 1\n                if pc.templ_tid not in used_tids:\n                    pcs.append(pc)\n                    used_tids.add(pc.templ_tid)\n                    break\n            inds[i] = ind\n        if pcs_len == len(pcs):\n            break  # no change\n        pcs = aggregate_pcs(pcs)\n    pcs = pcs[:npreds]\n\n    if DO_VALIDATE:\n        print_pcs('filtered', pcs)\n        N = len(query_idss_seq)\n        for pc in pcs:\n            assert pc.coords.shape == (N, 3), f'{pc.coords.shape} != ({N}, 3)'\n        c_pcs, c_dists = cluster_pcs(pcs, npreds, verbose=True)\n\n    return query_segments, [pc.coords for pc in pcs], pcs","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Protenix","metadata":{}},{"cell_type":"code","source":"def print_and_flush(message):\n    print(message)\n    sys.stdout.flush()\n\n\ndef get_protenix_configs(\n    n_cycle=PROTENIX_N_CYCLE, n_step=PROTENIX_N_STEP,\n    use_msa=True, use_rna_msa=True, use_template=True,\n    dump_dir=PROTENIX_OUT_DIR, model_name=PROTENIX_MODEL, dtype=PROTENIX_DTYPE,\n):\n    from configs.configs_base import configs as configs_base\n    from configs.configs_data import data_configs\n    from configs.configs_inference import inference_configs\n    from configs.configs_model_type import model_configs\n    from protenix.config.config import parse_configs\n\n    base = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n\n    def deep_update(t, p):\n        for k, v in p.items():\n            if isinstance(v, dict) and k in t and isinstance(t[k], dict):\n                deep_update(t[k], v)\n            else:\n                t[k] = v\n\n    deep_update(base, model_configs[model_name])\n    arg_str = \" \".join([\n        f\"--model_name {model_name}\",\n        f\"--dtype {dtype}\",\n        f\"--dump_dir {dump_dir}\",\n        f\"--use_msa {use_msa}\",\n        f\"--use_template {use_template}\",\n        f\"--use_rna_msa {use_rna_msa}\",\n        f\"--model.N_cycle {n_cycle}\",\n        f\"--sample_diffusion.N_step {n_step}\",\n    ])\n    cfg = parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)\n    cfg.data.template.prot_template_mmcif_dir = '/kaggle/input/stanford-rna-3d-folding-2/PDB_RNA'\n    cfg.infer_setting.dynamic_chunk_size = True\n    cfg.infer_setting.chunk_size_thresholds = {\n        \"512\": 256,\n        \"768\": 128,\n        \"896\": 64,\n        \"1024\": 32,\n        # > smallest_chunk_size: 32\n    }\n    # cfg.infer_setting.chunk_size = 32\n    # cfg.infer_setting.sample_diffusion_chunk_size = 1\n    return update_gpu_compatible_configs(cfg)\n\n\ndef extract_protenix_c1_coords(prediction, feat, chunk_seq_len):\n    raw_coords = prediction[\"coordinate\"]\n\n    if \"centre_atom_mask\" in feat:\n        mask = (feat[\"centre_atom_mask\"] == 1).to(raw_coords.device)\n    elif \"atom_to_tokatom_idx\" in feat:\n        m11 = (feat[\"atom_to_tokatom_idx\"] == 11).to(raw_coords.device)\n        m12 = (feat[\"atom_to_tokatom_idx\"] == 12).to(raw_coords.device)\n        c11, c12 = m11.sum(), m12.sum()\n        mask = m11 if abs(c11 - chunk_seq_len) < abs(c12 - chunk_seq_len) else m12\n    else:\n        mask = torch.zeros(raw_coords.shape[1], dtype=torch.bool, device=raw_coords.device)\n    \n    coords = raw_coords[:, mask, :].detach().cpu().numpy()\n    \n    # Collapse check\n    if coords.shape[1] > 1:\n        diffs = np.linalg.norm(coords[0, 1:] - coords[0, :-1], axis=-1)\n        if np.all(diffs < 1e-4):\n            print_and_flush(f\"WARNING! Collapsed coordinates detected!\")\n            return None\n    \n    if coords.shape[1] != chunk_seq_len:\n        if coords.shape[1] == 1 and chunk_seq_len > 1:\n            return None\n        padded = np.zeros((coords.shape[0], chunk_seq_len, 3), dtype=np.float32)\n        ml = min(coords.shape[1], chunk_seq_len)\n        padded[:, :ml, :] = coords[:, :ml, :]\n        coords = padded\n\n    return coords\n\n\ndef run_protenix(configs, runner, input_json_path, n_needed={}, call_id=''):\n    configs.input_json_path = input_json_path\n    dataset = InferenceDataset(configs)\n\n    pcs_dict = {}\n    for i in range(len(dataset)):\n        try:\n            data, atom_array, err = dataset[i]\n            sample_name = data.get(\"sample_name\", f\"sample_{i}\")\n            pcs_dict.setdefault(sample_name, [])\n            if err:\n                print_and_flush(f\"ERROR({call_id})! Protenix data error for {sample_name}: {err}!\")\n                coords_dict[sample_name] = None\n                continue\n            try:\n                sub_seq_len = data[\"N_token\"].item()  # roughly correct\n                configs = update_inference_configs(configs, sub_seq_len)\n                n_sample = n_needed.get(sample_name, 1)\n                configs.sample_diffusion.N_sample = n_sample\n                runner.update_model_configs(configs)\n\n                print_and_flush(f'Running protenix inference ({call_id}) for {sample_name}, n_sample: {n_sample}')\n                pred = runner.predict(data)\n                raw_coords = pred[\"coordinate\"]\n                coords = extract_protenix_c1_coords(\n                    pred, data[\"input_feature_dict\"], sub_seq_len,\n                )\n                if coords is None:\n                    continue\n                for i, conf in enumerate(pred['summary_confidence']):\n                    ranking_score, plddt, ptm, iptm = (\n                        float(conf[\"ranking_score\"]),\n                        float(conf[\"plddt\"]), float(conf[\"ptm\"]), float(conf[\"iptm\"]),\n                    )\n                    print_and_flush(f'... {sample_name}, call_id:{call_id}, i:{i}, ranking_score: {ranking_score:.4f}'\n                                    f' (plddt:{plddt:.2f}, ptm:{ptm:.2f}, iptm:{iptm:.2f})')\n                    scores = ([plddt, ranking_score] if PROTENIX_SORT_BY_PLDDT else [ranking_score, plddt]) + [ptm]\n                    pcs_dict[sample_name].append(PredCoords(*([sample_name] + scores + [coords[i, :, :]])))\n            except Exception as exc:\n                print(f\"ERROR({call_id})! Protenix inference failed for {sample_name}: {exc}!\")\n                import traceback; traceback.print_exc(); sys.stdout.flush()\n            finally:\n                gc.collect(); torch.cuda.empty_cache(); gc.collect()\n        except Exception as exc:\n            print_and_flush(f\"ERROR({call_id})! Protenix inference failed (i: {i}): {exc}!\")\n\n    return pcs_dict    \n\n\nclass ProtenixProcessor:\n    def __init__(\n        self, n_cycle=PROTENIX_N_CYCLE, n_step=PROTENIX_N_STEP,\n        use_msa=True, use_rna_msa=True, use_template=True,\n        dump_dir=PROTENIX_OUT_DIR, model_name=PROTENIX_MODEL,\n    ):\n        print(f'Starting ProtenixProcessor, n_cycle:{n_cycle}, n_step:{n_step}'\n              f', use_msa:{use_msa}, use_rna_msa:{use_rna_msa}, use_template:{use_template}')\n        self.configs = get_protenix_configs(\n            n_cycle=n_cycle, n_step=n_step, dump_dir=dump_dir, model_name=model_name,\n            use_msa=use_msa, use_rna_msa=use_rna_msa, use_template=use_template,\n        )\n        self.runner = InferenceRunner(self.configs)\n\n    def run(self, input_json_path, n_needed={}, call_id=''):\n        return run_protenix(self.configs, self.runner,\n                            input_json_path=input_json_path,\n                            n_needed=n_needed, call_id=call_id)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def _get_json_entry(prev_seq, prev_count, msa_file, use_msa_file, call_id, target_id):\n    entry = {\n        \"sequence\": prev_seq,\n        \"count\": prev_count,\n    }\n    if use_msa_file:\n        if os.path.isfile(msa_file):\n            entry[\"unpairedMsaPath\"] = msa_file\n        else:\n            print_and_flush(f'WARNING({call_id})!'\n                            f' MSA file is not found for {target_id}: {msa_file}!')\n    return entry\n\n\ndef are_files_equal(file1_path, file2_path):\n    with open(file1_path, 'r') as f1, open(file2_path, 'r') as f2:\n        return f1.read() == f2.read()\n\n\ndef prepare_protenix_json(target_id, sequence, msa_file,\n                          call_id='', query_db=None, ligand_smiles=None, query_segments=None,\n                          max_length=PROTENIX_MAX_LEN, output_path=PROTENIX_OUT_DIR,\n                          split_chains=PROTENIX_JSON_SPLIT_CHAINS):\n    if msa_file is not None:\n        msa_file = str(msa_file)\n    if len(sequence) <= max_length:\n        use_msa_file = msa_file is not None and 0 < len(msa_file)\n        seq = sequence\n        if not split_chains:\n            query_segments = None\n    else:\n        print_and_flush(f\"WARNING({call_id})! Sequence too long ({len(sequence)} > {max_length})\"\n                        f\", truncating for {target_id}, ignoring MSA file.\")\n        use_msa_file = False\n        query_segments = None\n        seq = sequence[:max_length]\n    if query_segments is None:\n        query_segments = [(0, len(seq), 'A')]\n\n    if len(query_segments) == 1:\n        msa_files = [msa_file]\n    elif not use_msa_file:\n        msa_files = [None] * len(query_segments)\n    else:\n        msa_dir = Path(output_path) / 'MSA'\n        msa_dir.mkdir(parents=True, exist_ok=True)\n        msa_file_pre = '.'.join(os.path.basename(msa_file).split('.')[:-2])\n        msa_files = []\n        prev_msa = None\n        for i, (b, e, chain_id) in enumerate(query_segments):\n            msa = str(msa_dir / f'{msa_file_pre}_{call_id}_{i}.MSA.fasta')\n            is_msa, _ = truncate_msa_file(msa_file, b, e, chain_id, output_msa_file=msa)\n            if prev_msa and is_msa and are_files_equal(prev_msa, msa):\n                msa_files.append(prev_msa)\n            else:\n                if not is_msa:\n                    msa = ''\n                msa_files.append(msa)\n                prev_msa = msa\n\n    sequences_content = []\n    prev_seq = None\n    prev_count = 0\n    prev_msa = None\n    for i, (b, e, chain_id) in enumerate(query_segments):\n        chain_seq = seq[b:e]\n        chain_msa = msa_files[i]\n        if chain_seq == prev_seq and chain_msa == prev_msa:\n            prev_count += 1\n        else:\n            if prev_seq is not None:\n                use_prev_msa = prev_msa is not None and len(prev_msa) > 0\n                sequences_content.append({\"rnaSequence\": _get_json_entry(\n                    prev_seq, prev_count, prev_msa, use_prev_msa, call_id, target_id,\n                )})\n            prev_seq = chain_seq\n            prev_msa = chain_msa\n            prev_count = 1\n\n    if prev_seq is not None:\n        use_prev_msa = prev_msa is not None and len(prev_msa) > 0\n        sequences_content.append({\"rnaSequence\": _get_json_entry(\n            prev_seq, prev_count, prev_msa, use_prev_msa, call_id, target_id,\n        )})\n\n    if isinstance(ligand_smiles, str):\n        sequences_content.append({\n            \"ligand\": {\n                \"ligand\": ligand_smiles,\n                \"count\": 1,\n            },\n        })\n\n    if isinstance(ligand_smiles, str):\n        sequences_content.append({\n            \"ligand\": {\n                \"ligand\": ligand_smiles,\n                \"count\": 1,\n            },\n        })\n\n    input_json = {\n        \"name\": target_id,\n        \"sequences\": sequences_content,\n    }\n    json_path = Path(output_path) / \"input_json\" / f\"{call_id}_{target_id}.json\"\n    json_path.parent.mkdir(parents=True, exist_ok=True)\n    with open(json_path, \"w\") as f:\n        json.dump([input_json], f, indent=4)\n\n    return json_path\n\n\ndef run_protenix_inference(processor, target_id, query_seq, seed, n_needed, msa_file, call_id='',\n                           max_length=PROTENIX_MAX_LEN, output_path=PROTENIX_OUT_DIR,\n                           query_segments=None, query_db=None, ligand_smiles=None):\n    if T_MAX is not None and time.time() > T_MAX:\n        return []\n    input_json_path = prepare_protenix_json(target_id, query_seq,\n                                            msa_file=msa_file,\n                                            call_id=call_id,\n                                            query_db=query_db,\n                                            query_segments=query_segments,\n                                            ligand_smiles=ligand_smiles,\n                                            max_length=max_length,\n                                            output_path=output_path)\n    processor.configs.seed = seed\n    processor.configs.seeds = [seed]\n    seed_everything(seed, deterministic=True)\n    pcs_dict = processor.run(input_json_path, call_id=call_id,\n                             n_needed={target_id: n_needed})\n    pcs = pcs_dict.get(target_id)\n    return [] if pcs is None else pcs\n\n\ndef _get_protenix_predictions(\n    gpu_id, seed, protenix_queue, msa_path, chunk_size, chunk_ramp,\n    use_unpaired_msa_file, use_trunc_msa_files, min_n_sample,\n    use_msa, use_rna_msa, use_template, use_ligand,\n    max_length, output_path, order_multi_samples,\n):\n    t0 = time.time()\n    os.environ[\"CUDA_VISIBLE_DEVICES\"] = str(gpu_id)\n    processor = ProtenixProcessor(\n        use_msa=use_msa, use_rna_msa=use_rna_msa, use_template=use_template,\n    )\n    if not use_msa:\n        use_unpaired_msa_file = False\n        use_trunc_msa_files = False\n    print_and_flush(f'Protenix model is loaded in {time.time()-t0:.1f} seconds on GPU:{gpu_id}.')\n\n    pcs = {}\n    info = {}\n    iter_times = {}\n    t_single_queries = 0.0\n    t_multi_queries = 0.0\n    is_time_up = False\n    for sweep in range(2):\n        if is_time_up:\n            break\n        for idx, target_id, query_seq, segments, n_needed, query_db, ligand_smiles, align_coords in protenix_queue:\n            if not use_ligand:\n                ligand_smiles = None\n            msa_file = msa_path / f\"{target_id}.MSA.fasta\" if use_unpaired_msa_file else None\n            ncalls = 0\n            call_id = f'{gpu_id}-{ncalls}'\n            is_single_query = (len(query_seq) <= max_length)\n            query_maxiters = PROTENIX_QUERY_MAXITERS[is_single_query]\n            if PROTENIX_AS_NEEDED_MAXITERS:\n                query_maxiters = max(query_maxiters, n_needed)\n            query_maxseconds = PROTENIX_QUERY_MAXSECONDS[is_single_query]\n            print_and_flush(f'\\nProtenix queue ({call_id}), idx:{idx}, tid:{target_id}, single:{is_single_query}'\n                            f', maxiters:{query_maxiters}, maxseconds:{query_maxseconds}')\n            t_start = time.time()\n            pcs.setdefault(idx, [])\n            try:\n                for niter in range(query_maxiters):\n                    print_and_flush(f'Protenix iteration ({call_id}):{target_id}'\n                                    f', niter:{niter}, time:{time.time()-t_start:.1f}s')\n                    t_iter = time.time()\n                    niter_needed = (\n                        min_n_sample if PROTENIX_AS_NEEDED_MAXITERS\n                        else (max(min_n_sample, n_needed) if niter == 0 and sweep == 0 else min_n_sample)\n                    )\n                    if is_single_query:\n                        ncalls += 1\n                        call_id = f'{gpu_id}-{ncalls}'\n                        print_and_flush(f'Calling protenix inference ({call_id}) for {target_id}'\n                                        f', len(query_seq): {len(query_seq)}')\n                        pcs[idx] += run_protenix_inference(\n                            processor, seed=seed+10001*niter, call_id=call_id,\n                            target_id=target_id, query_seq=query_seq, \n                            n_needed=niter_needed, msa_file=msa_file,\n                            query_segments=segments, query_db=query_db, ligand_smiles=ligand_smiles,\n                            max_length=max_length, output_path=output_path)\n                    else:\n                        seen_chain_ids = set()\n                        tid_pcs = [PredCoords(target_id, 0.0, 0.0, 0.0,\n                                              np.zeros((len(query_seq), 3)) * np.nan)] * niter_needed\n                        tid_n = 0\n                        for j, (b, e, chain_id) in enumerate(segments):\n                            if PROTENIX_MULTI_CHAIN_SKIP and chain_id in seen_chain_ids:\n                                continue\n                            seen_chain_ids.add(chain_id)\n                            pred_ranges = [(b, min(e, b + max_length))]\n                            if chunk_size is not None:\n                                assert chunk_size < max_length\n                                while pred_ranges[-1][1] < e:\n                                    next_b = pred_ranges[-1][1] - chunk_size\n                                    pred_ranges.append((next_b, min(e, next_b + max_length)))\n                            print_and_flush(f'Truncated ranges ({gpu_id}:{ncalls}) for {target_id}: {pred_ranges}')\n                            is_first = True\n                            for pi, (rb, re) in enumerate(pred_ranges):\n                                ncalls += 1\n                                if use_unpaired_msa_file and use_trunc_msa_files:\n                                    trunc_msa_file = output_path / f'{ncalls:02d}_{target_id}_{chain_id}_{j}-{pi}.MSA.fasta'\n                                    if not truncate_msa_file(msa_file, rb, re, chain_id,\n                                                             output_msa_file=trunc_msa_file)[0]:\n                                        trunc_msa_file = None\n                                else:\n                                    trunc_msa_file = None\n                                call_id = f'{gpu_id}-{ncalls}'\n                                print_and_flush(f'Calling protenix inference ({call_id}) for {target_id}:'\n                                                f'j:{j},pi:{pi}|chain:{chain_id}[{rb}:{re}]')\n                                new_pcs = run_protenix_inference(\n                                    processor, seed=seed+10001*niter+j*107+pi, call_id=call_id,\n                                    target_id=target_id, query_seq=query_seq[rb:re],\n                                    n_needed=niter_needed, msa_file=trunc_msa_file,\n                                    max_length=max_length, output_path=output_path,\n                                )\n                                if order_multi_samples:\n                                    print(f'Sorting multi query ({call_id}) for {target_id}.')\n                                    new_pcs.sort(key=get_pc_score, reverse=True)\n                                tid_n += 1\n                                if is_first:\n                                    is_first = False\n                                    for k, (_, new_score, new_plddt, new_ptm, new_coords) in enumerate(new_pcs):\n                                        tid, score, plddt, ptm, coords = tid_pcs[k]\n                                        coords[rb:re, :] = new_coords\n                                        tid_pcs[k] = PredCoords(tid, score+new_score, plddt+new_plddt,\n                                                                ptm+new_ptm, coords)\n                                else:\n                                    ramp_weights = np.linspace(0.0, 1.0, chunk_size)[:, None]\n                                    for k, (_, new_score, new_plddt, new_ptm, new_coords) in enumerate(new_pcs):\n                                        tid, score, plddt, ptm, coords = tid_pcs[k]\n                                        R, t = kabsch(new_coords[:chunk_size, :], coords[rb:rb+chunk_size, :])\n                                        new_coords = new_coords @ R.T + t\n                                        if chunk_ramp:\n                                            coords[rb:rb+chunk_size, :] *= (1.0 - ramp_weights)\n                                            coords[rb:rb+chunk_size, :] += ramp_weights * new_coords[:chunk_size, :]\n                                        coords[rb+chunk_size:re, :] = new_coords[chunk_size:, :]\n                                        tid_pcs[k] = PredCoords(tid, score+new_score, plddt+new_plddt,\n                                                                ptm+new_ptm, coords)\n        \n                        for k, (tid, score, plddt, ptm, coords) in enumerate(tid_pcs):\n                            tid_pcs[k] = PredCoords(tid, score/tid_n, plddt/tid_n, ptm/tid_n, coords)\n    \n                        if align_coords is not None:\n                            print(f'Orienting multi query ({call_id}) to provided coordinates.')\n                            for k, (tid, score, plddt, ptm, coords) in enumerate(tid_pcs):\n                                for (b, e, _) in segments:\n                                    R, t = kabsch(coords[b:e, :], align_coords[b:e, :])\n                                    coords[b:e, :] = coords[b:e, :] @ R.T + t\n                                tid_pcs[k] = PredCoords(tid, score, plddt, ptm, coords)\n                        elif PROTENIX_ORIENT_CHUNKS:\n                            ncalls += 1\n                            chunk_len = int(max_length / len(segments))\n                            call_id = f'{gpu_id}-{ncalls}'\n                            print_and_flush(f'Orienting chunks ({call_id}) for {target_id}'\n                                            f', nsegments: {len(segments)}, chunk_len: {chunk_len}')\n                            trunc_segments = []\n                            for (b, e, chain_id) in segments:\n                                mid = (b + min(b+max_length, e)) // 2\n                                start = max(b, mid - chunk_len//2)\n                                trunc_segments.append((start, min(e, start + chunk_len), chain_id))\n                            if use_unpaired_msa_file:\n                                trunc_msa_file = output_path / f'{ncalls:02d}_{target_id}_orc.MSA.fasta'\n                                if not all_chain_truncate_msa_file(msa_file, trunc_segments,\n                                                                   output_msa_file=trunc_msa_file)[0]:\n                                    trunc_msa_file = None\n                            else:\n                                trunc_msa_file = None\n                            seq = ''.join([query_seq[tb:te] for (tb, te, _) in trunc_segments])\n                            print_and_flush(f'Calling protenix inference ({call_id}) for {target_id}'\n                                            f', len(seq): {len(seq)}, chunk_len: {chunk_len}')\n                            orc_segments = []\n                            pos = 0\n                            for tb, te, chain_id in trunc_segments:\n                                new_pos = pos + te - tb\n                                orc_segments.append((pos, new_pos, chain_id))\n                                pos = new_pos\n                            assert len(seq) == pos, f'len(seq): {len(seq)}, pos: {pos}'\n                            orc_pcs = run_protenix_inference(processor, seed=seed+10001*niter+len(segments),\n                                                             call_id=call_id, target_id=target_id,\n                                                             query_seq=seq, query_segments=orc_segments,\n                                                             n_needed=niter_needed, msa_file=trunc_msa_file,\n                                                             max_length=max_length, output_path=output_path)\n                            for k, (_a, _b, _c, _d, orc_coords) in enumerate(orc_pcs):\n                                tid, score, plddt, ptm, coords = tid_pcs[k]\n                                for (b, e, _), (tb, te, _), (ob, oe, _) in zip(segments, trunc_segments, orc_segments):\n                                    R, t = kabsch(coords[tb:te, :], orc_coords[ob:oe, :])\n                                    coords[b:e, :] = coords[b:e, :] @ R.T + t\n                                tid_pcs[k] = PredCoords(tid, score, plddt, ptm, coords)\n        \n                        pcs[idx] += tid_pcs\n\n                    t_now = time.time()\n                    iter_time = (t_now - t_iter) / (niter_needed / min_n_sample)\n                    sum_time, count = iter_times.get(idx, (0.0, 0))\n                    sum_time += iter_time\n                    count += 1\n                    iter_times[idx] = (sum_time, count)\n                    avg_iter_time = sum_time / count\n                    if t_now - t_start + avg_iter_time > query_maxseconds:\n                        break\n                    if T_MAX is not None and t_now + avg_iter_time > T_MAX:\n                        is_time_up = True\n                        break\n            except Exception as exception:\n                print_and_flush(f\"ERROR({call_id})! Protenix failed for {target_id}[{idx}]: {exception}!\")\n            except RuntimeError as rexception:\n                print_and_flush(f\"RuntimeERROR({call_id})! Protenix failed for {target_id}[{idx}]: {rexception}!\")\n            finally:\n                torch.cuda.empty_cache()\n                gc.collect()\n    \n            t_end = time.time()\n            t_seconds = t_end - t_start\n            if is_single_query:\n                t_single_queries += t_seconds\n            else:\n                t_multi_queries += t_seconds\n            info[idx] = (target_id, ncalls, t_seconds)\n            print_and_flush(f'Finished protenix inference for {target_id}'\n                            f' on GPU:{gpu_id} in {ncalls} calls and {t_seconds:.1f} seconds'\n                            f' (len:{len(pcs[idx])}).')\n            if T_MAX is not None and t_end > T_MAX:\n                is_time_up = True\n                break\n\n    for idx in pcs.keys():\n        pcs[idx] = [PredCoords(f'{pc.templ_tid}:{gpu_id}:{i}', pc.filled_score, pc.seqmr, pc.score, pc.coords)\n                    for i, pc in enumerate(pcs[idx])]\n        pcs[idx] = sorted(pcs[idx], key=get_pc_score, reverse=True)\n\n    print('\\nProtenix scores:')\n    t_total = 0.0\n    for idx, (target_id, ncalls, t_seconds) in sorted(info.items()):\n        t_total += t_seconds\n        print(f'idx:{idx}, npcs:{len(pcs[idx])}, gpu:{gpu_id}, tid:{target_id}, ncalls:{ncalls}, time:{t_seconds:.1f}s')\n        print_pcs('    ', pcs[idx], do_print_score=True)\n    print(f'\\nTotal protenix time (gpu:{gpu_id}): {t_total:.1f}s')\n    print(f'    single queries time: {t_single_queries:.1f}s')\n    print(f'     multi queries time: {t_multi_queries:.1f}s')\n    print('')\n    sys.stdout.flush()\n\n    return pcs\n\n\ndef get_protenix_predictions(protenix_queue,\n                             msa_path=MSA_DIR,\n                             chunk_size=PROTENIX_CHUNK_SIZE,\n                             chunk_ramp=PROTENIX_CHUNK_RAMP,\n                             use_msas=PROTENIX_USE_MSAS,\n                             use_rna_msas=PROTENIX_USE_RNA_MSAS,\n                             use_templates=PROTENIX_USE_TEMPLATES,\n                             use_ligand=PROTENIX_USE_LIGAND,\n                             use_unpaired_msa_file=PROTENIX_USE_UNPAIRED_MSA_FILES,\n                             use_trunc_msa_files=PROTENIX_USE_TRUNC_MSA_FILES,\n                             min_n_sample=PROTENIX_MIN_N_SAMPLE,\n                             max_length=PROTENIX_MAX_LEN,\n                             output_path=PROTENIX_OUT_DIR,\n                             order_multi_samples=PROTENIX_ORDER_MULTI_SAMPLES,\n                             seed=PROTENIX_SEED):\n    gpu_ids = list(range(torch.cuda.device_count()))\n    if not gpu_ids:\n        print(f'WARNING! No available GPU, skipping Protenix!')\n        return {}\n\n    print(f'Available GPUs: {gpu_ids}')\n    seeds = [seed + i*7657 for i in range(len(gpu_ids))]\n    print(f'Seeds: {seeds}')\n\n    all_pcs = Parallel(n_jobs=len(gpu_ids), backend=\"loky\")(\n        delayed(_get_protenix_predictions)(\n            gpu_id_, seed_, protenix_queue, msa_path, chunk_size, chunk_ramp,\n            use_unpaired_msa_file, use_trunc_msa_files, min_n_sample,\n            use_msa_, use_rna_msa_, use_template_, use_ligand,\n            max_length, output_path, order_multi_samples,\n        )\n        for (gpu_id_, seed_, use_msa_, use_rna_msa_, use_template_) in zip(\n            gpu_ids, seeds, use_msas, use_rna_msas, use_templates,\n        )\n    )\n    pcs = {}\n    for gpu_pcs in all_pcs:\n        for idx, idx_pcs in gpu_pcs.items():\n            if idx not in pcs:\n                pcs[idx] = idx_pcs\n            else:\n                pcs[idx] += idx_pcs\n\n    for idx in pcs.keys():\n        pcs[idx] = sorted(pcs[idx], key=get_pc_score, reverse=True)\n\n    # idx, target_id, query_seq, segments, n_needed, query_db, ligand_smiles, align_coords = protenix_queue\n    queue_data = {rec[0]: (rec[1], rec[3], rec[4]) for rec in protenix_queue}\n    coords_dict = {}\n    for idx in sorted(pcs.keys()):\n        tid, segments, n_needed = queue_data[idx]\n        pcs[idx] = aggregate_pcs(pcs[idx])\n        pcs[idx] = fill_pcs_coords(pcs[idx], segments)\n        if DO_VALIDATE:\n            print(f'idx: {idx}, tid: {tid}, npcs: {len(pcs[idx])}')\n            c_pcs, c_dists = cluster_pcs(pcs[idx], len(pcs[idx]), verbose=True)\n            sys.stdout.flush()\n        coords_dict[idx] = [pc.coords for pc in pcs[idx][:n_needed]]\n    return coords_dict\n\n\ndef get_protenix_n_needed(pcs, npreds):\n    if len(pcs) > npreds - len(USE_PROTENIX_SCORES):\n        i = npreds - len(USE_PROTENIX_SCORES)\n        new_pcs = pcs[:i]\n        for score_limit in USE_PROTENIX_SCORES:\n            if i >= len(pcs):\n                break\n            pc = pcs[i]\n            if get_pc_score(pc)[0] >= score_limit:\n                new_pcs.append(pc)\n            i += 1\n        pcs = new_pcs\n    return max(0, npreds - len(pcs))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## TBM + Protenix","metadata":{}},{"cell_type":"code","source":"npreds = NUM_PREDS\nT_PRED_START = time.time()\n\nresults = []\nprotenix_queue = []\nall_pcs = {}  # used only for validation\nfor idx, (row_id, row) in enumerate(test_seqs.iterrows()):\n    target_id = row[\"target_id\"]\n    print(f'\\n>>> {target_id} ({idx+1}/{len(test_seqs)})')\n    temporal_cutoff = row[\"temporal_cutoff\"] if USE_TEMPORAL_CUTOFF else None\n    query_segments, coords, pcs = get_tbm_predictions(\n        row_id, npreds=npreds, temporal_cutoff=temporal_cutoff,\n    )\n    all_pcs[target_id] = pcs\n    query_seq = row[\"sequence\"]\n    rnastruct_row = rnastruct_test.loc[row_id, :]\n    query_db = ''.join(get_dotbrackets(rnastruct_row))\n    results.append((target_id, query_seq, query_segments, query_db, coords))\n\n    if USE_PROTENIX:\n        if DO_VALIDATE and target_id in SKIP_PROTENIX_TIDS:\n            print(f'Skipping Protenix for {target_id}.')\n        else:\n            n_needed = get_protenix_n_needed(pcs, npreds)\n            if n_needed > 0:\n                align_coords = coords[0] if PROTENIX_MULTI_ALIGN_TO_BEST_TBM and len(coords) > 0 else None\n                protenix_queue.append((idx, target_id, query_seq, query_segments,\n                                       n_needed, query_db,\n                                       row[\"ligand_SMILES\"], align_coords))\nT_TBM = time.time() - T_PRED_START\n\nprint('')\nT_PROTENIX_START = time.time()\nif USE_PROTENIX:\n    print('PROTENIX begin')\n    protenix_queue.sort(key=lambda x: len(x[2]), reverse=False)\n    for idx, protenix_coords in get_protenix_predictions(protenix_queue).items():\n        target_id, query_seq, query_segments, query_db, coords = results[idx]\n        n_keep = npreds - len(protenix_coords)\n        all_pcs[target_id] = all_pcs[target_id][:n_keep]\n        coords = coords[:n_keep] + protenix_coords\n        results[idx] = (target_id, query_seq, query_segments, query_db, coords)\n    print('PROTENIX end\\n')\nT_PROTENIX = time.time() - T_PROTENIX_START\n\nT_POSTPROCESS_START = time.time()\nif USE_COORD_POSTPROCESSING:\n    for idx, (target_id, query_seq, query_segments, query_db, coords) in tqdm(enumerate(results),\n                                                                              total=len(results),\n                                                                              desc=\"Postprocessing\"):\n        results[idx] = (target_id, query_seq, query_segments, query_db,\n                        [refine_rna_coords(query_seq, query_segments, query_db, c)\n                         for c in coords])\nT_POSTPROCESS = time.time() - T_POSTPROCESS_START\n\nfor idx, (target_id, query_seq, query_segments, query_db, coords) in enumerate(results):\n    if len(coords) > npreds:\n        print(f'WARNING: too many predictions for {target_id}: {len(coords)}!')\n        coords = coords[:npreds]\n    query_segments_wo_ids = [(b, e) for (b, e, _) in query_segments]\n    n_needed = npreds - len(coords)\n    if n_needed > 0:\n        print(f'Missing de novo coords fill in for {target_id}: {n_needed}!')\n        coords += [np.cumsum(np.ones((len(query_seq), 3), dtype=float), axis=0)] * n_needed\n    results[idx] = (target_id, query_seq, query_segments, query_db, coords)\n\npreds = []\nfor target_id, query_seq, query_segments, query_db, coords in results:\n    for j in range(len(query_seq)):\n        pred_row = {'ID': f'{target_id}_{j+1}', 'resname': query_seq[j], 'resid': j+1}\n        for i in range(NUM_PREDS):\n            pred_row[f\"x_{i+1}\"] = coords[i][j][0]\n            pred_row[f\"y_{i+1}\"] = coords[i][j][1]\n            pred_row[f\"z_{i+1}\"] = coords[i][j][2]\n        preds.append(pred_row)\nsubmission_df = pd.DataFrame(preds)\n\n# Ensure the submission file has the correct format\ncolumn_order = [\"ID\", \"resname\", \"resid\"]\nfor i in range(1, 1+NUM_PREDS):\n    for coord in [\"x\", \"y\", \"z\"]:\n        column_order.append(f\"{coord}_{i}\")\nsubmission_df = submission_df[column_order]\n\n# Clip explicitly (competition clips coords; prevent explosions)\ncoord_cols = [c for c in submission_df.columns if c.startswith((\"x_\", \"y_\", \"z_\"))]\nif np.isnan(submission_df[coord_cols].values).any():\n    print(\"\\n\\n!!! NaNs in submission_df !!!\\n\\n\")\nsubmission_df[coord_cols] = np.round(submission_df[coord_cols], decimals=3)\nsubmission_df[coord_cols] = submission_df[coord_cols].fillna(0.0).clip(COORDS_MIN, COORDS_MAX)\n\n# Save the submission file\nsubmission_df.to_csv(\"submission.csv\", index=False)\nprint(f\"Saved predictions for {len(test_seqs)} RNA sequences to submission.csv.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f'TBM time         : {T_TBM:.1f}s')\nprint(f'Protenix time    : {T_PROTENIX:.1f}s')\nprint(f'Postprocess time : {T_POSTPROCESS:.1f}s')\nT = time.time()\nprint(f'Prediction time  : {T - T_PRED_START:.1f}s')\nprint(f'Notebook time    : {T - T_NOTEBOOK_START:.1f}s')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Validation","metadata":{}},{"cell_type":"code","source":"%%time\n\nmetric_path = \"/kaggle/input/datasets/gabalz/rnastruct/metric.py\"\nusalign_path = \"/kaggle/input/datasets/metric/usalign/USalign\"\n\nskip_target_ids = SKIPPED_VALID_TARGET_IDS\nscore_ref = {\n    '8ZNQ': [0.232, 0.186],\n    '9CFN': [0.373, 0.160],\n    '9E74': [0.664, 0.649],\n    '9E75': [0.768, 0.740],\n    '9E9Q': [0.560, 0.560],\n    '9EBP': [0.531, 0.513],\n    '9G4J': [0.922, 0.922],\n    '9G4P': [0.190, 0.106],\n    '9G4Q': [0.227, 0.193],\n    '9G4R': [0.376, 0.256],\n    '9HRO': [0.292, 0.285],\n    '9I9W': [0.291, 0.279],\n    '9IWF': [0.679, 0.609],\n    '9J09': [0.169, 0.136],\n    '9JFO': [0.183, 0.183],\n    '9JFS': [0.134, 0.134],\n    '9JGM': [0.301, 0.301],\n    '9KGG': [0.696, 0.636],\n    '9LEC': [0.856, 0.856],\n    '9LEL': [0.690, 0.690],\n    '9LJN': [0.702, 0.702],\n    '9MME': [0.850, 0.822],\n    '9OBM': [0.353, 0.244],\n    '9OD4': [0.382, 0.251],\n    '9QZJ': [0.206, 0.178],\n    '9RVP': [0.296, 0.169],\n    '9WHV': [0.244, 0.244],\n    '9ZCC': [0.389, 0.389],\n}\n\ndef scoring_worker(args):\n    target_id, group_native, group_pred = args\n    t0 = time.time()\n\n    # load score inside each worker (safe for multiprocessing)\n    module_globals = runpy.run_path(metric_path)\n    score_func = module_globals[\"score\"]\n\n    final_score, all_scores = score_func(group_native, group_pred, \"ID\",\n                                         return_all_scores=True)\n    print(f'... {target_id}: '\n          + ', '.join([f'{idx}|{s:.3f}' for _, idx, s in all_scores])\n          + f'    [{time.time()-t0:.1f}s]')\n    return target_id, final_score, all_scores\n\n\ndef get_max_scores(all_scores, all_pcs):\n    max_scores = {}\n    for tid, idx, score in all_scores:\n        pcs = all_pcs[tid]\n        pc = pcs[idx-1] if idx <= len(pcs) else None\n        max_score = max_scores.get(idx, (-np.inf, None))[0]\n        if max_score < score:\n            max_scores[idx] = (score, pc)\n    return [(r[1].templ_tid if r[1] is not None else '',\n             f'{r[1].filled_score:.3f}' if r[1] is not None else '',\n             i, r[0]) for i, r in sorted(max_scores.items())]\n\n\nif DO_VALIDATE and valid_labels is not None:\n    shutil.copy2(usalign_path, '/kaggle/working/USalign')\n    os.chmod('/kaggle/working/USalign', 0o755)\n\n    sol = valid_labels.copy()\n    sub = pd.read_csv(\"/kaggle/working/submission.csv\")\n    assert not np.isnan(sub[coord_cols].values).any()\n    # target_id is ID without the residue suffix\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    tasks = []\n    for target_id, group_native in sol.groupby(\"target_id\"):\n        if target_id in skip_target_ids:\n            continue\n        group_pred = sub[sub[\"target_id\"] == target_id]\n        if len(group_pred) == 0:\n            continue\n        tasks.append((target_id, group_native, group_pred))\n\n    results = []\n    with ProcessPoolExecutor(max_workers=NWORKERS) as executor:\n        futures = [executor.submit(scoring_worker, t) for t in tasks]\n\n        for i, future in enumerate(as_completed(futures), 1):\n            results.append(future.result())\n\n    print('')\n    result_scores = []\n    for i, (target_id, score, all_scores) in enumerate(sorted(results)):\n        result_scores.append(score)\n        row = test_seqs[test_seqs['target_id'] == target_id].iloc[0, :]\n        query_seq = row['sequence']\n        query_segments = get_chain_segments(row)\n        extra = f', len:{len(query_seq)}'\n        if len(query_segments) > 1:\n            extra += f'|{\",\".join([f\"{cid}:{e-b}\" for b, e, cid in query_segments])}'\n        print(f'{i+1:2d}, {target_id}({len(query_segments)}), {score:.3f}'\n              f',   {score-score_ref[target_id][0]:6.3f}'\n              f', {score-score_ref[target_id][1]:6.3f}'\n              f'{extra}')\n    print('----')\n    mean_score = np.mean(result_scores)\n    ref1_mean_score = np.mean([score[0] for tid, score in score_ref.items() if tid not in skip_target_ids])\n    ref2_mean_score = np.mean([score[1] for tid, score in score_ref.items() if tid not in skip_target_ids])\n    print(f'Mean score (n={len(result_scores)}): {mean_score:.3f}'\n          f', ref1_diff: {mean_score-ref1_mean_score:.3f}'\n          f', ref2_diff: {mean_score-ref2_mean_score:.3f}')\n    print('----')\n    for i, (target_id, score, all_scores) in enumerate(sorted(results)):\n        print(f'{i+1:2d}, {target_id}({len(query_segments)}), {score:.3f}:  '\n              + ', '.join([f'{tt}|{ts}|{i}|{s:.3f}'\n                           for tt, ts, i, s in get_max_scores(all_scores, all_pcs)]))\nprint('')","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}