{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.11"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":87793,"databundleVersionId":12255112,"sourceType":"competition"},{"sourceId":395332,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":324795,"modelId":345618},{"sourceId":395352,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":324804,"modelId":345630}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":1724.203471,"end_time":"2025-05-17T08:30:26.446107","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-05-17T08:01:42.242636","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"c5a00ab6","cell_type":"code","source":"!pip install /kaggle/input/spektral/pytorch/default/1/spektral-1.3.1-py3-none-any.whl > /dev/null 2>&1\n!pip install /kaggle/input/bio/pytorch/default/1/biopython-1.85-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl > /dev/null 2>&1","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:01:46.882638Z","iopub.status.busy":"2025-05-17T08:01:46.882168Z","iopub.status.idle":"2025-05-17T08:01:55.432617Z","shell.execute_reply":"2025-05-17T08:01:55.431607Z"},"papermill":{"duration":8.559175,"end_time":"2025-05-17T08:01:55.434194","exception":false,"start_time":"2025-05-17T08:01:46.875019","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"4570bbce","cell_type":"code","source":"import os\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '2'","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:01:55.445400Z","iopub.status.busy":"2025-05-17T08:01:55.445163Z","iopub.status.idle":"2025-05-17T08:01:55.448961Z","shell.execute_reply":"2025-05-17T08:01:55.448450Z"},"papermill":{"duration":0.010479,"end_time":"2025-05-17T08:01:55.449963","exception":false,"start_time":"2025-05-17T08:01:55.439484","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"76239150","cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport tensorflow as tf\nfrom tensorflow.keras.layers import Input, Dense, Layer, Dropout, BatchNormalization\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.callbacks import EarlyStopping, ReduceLROnPlateau, Callback\nfrom tensorflow.keras.mixed_precision import set_global_policy\nfrom tensorflow.keras.optimizers import AdamW\nimport spektral.layers as gnn_layers\nfrom sklearn.preprocessing import StandardScaler\nfrom concurrent.futures import ProcessPoolExecutor\nimport multiprocessing\nimport gc\nimport pickle\nfrom Bio import SeqIO\nimport subprocess\nimport warnings\nimport logging\n\n\nwarnings.filterwarnings('ignore')\nset_global_policy('mixed_float16')","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2025-05-17T08:01:55.460563Z","iopub.status.busy":"2025-05-17T08:01:55.460025Z","iopub.status.idle":"2025-05-17T08:02:13.741570Z","shell.execute_reply":"2025-05-17T08:02:13.740724Z"},"papermill":{"duration":18.288017,"end_time":"2025-05-17T08:02:13.742971","exception":false,"start_time":"2025-05-17T08:01:55.454954","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"7424c868","cell_type":"code","source":"# Set up logging\nlogging.basicConfig(filename='submission.log', level=logging.INFO, \n                    format='%(asctime)s - %(levelname)s - %(message)s')","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:02:13.754508Z","iopub.status.busy":"2025-05-17T08:02:13.753979Z","iopub.status.idle":"2025-05-17T08:02:13.757650Z","shell.execute_reply":"2025-05-17T08:02:13.757066Z"},"papermill":{"duration":0.01036,"end_time":"2025-05-17T08:02:13.758662","exception":false,"start_time":"2025-05-17T08:02:13.748302","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"7d6bf776","cell_type":"code","source":"# Set seeds\nnp.random.seed(42)\ntf.random.set_seed(42)","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:02:13.768816Z","iopub.status.busy":"2025-05-17T08:02:13.768612Z","iopub.status.idle":"2025-05-17T08:02:13.771825Z","shell.execute_reply":"2025-05-17T08:02:13.771328Z"},"papermill":{"duration":0.00929,"end_time":"2025-05-17T08:02:13.772804","exception":false,"start_time":"2025-05-17T08:02:13.763514","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"a77140f4","cell_type":"code","source":"# ==============================\n# GPU Configuration\n# ==============================\n\nphysical_devices = tf.config.list_physical_devices('GPU')\nfor device in physical_devices:\n    tf.config.experimental.set_memory_growth(device, True)","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:02:13.782683Z","iopub.status.busy":"2025-05-17T08:02:13.782472Z","iopub.status.idle":"2025-05-17T08:02:14.894305Z","shell.execute_reply":"2025-05-17T08:02:14.893587Z"},"papermill":{"duration":1.117962,"end_time":"2025-05-17T08:02:14.895440","exception":false,"start_time":"2025-05-17T08:02:13.777478","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"32da0baf","cell_type":"code","source":"class ReshapeLayer(Layer):\n    def __init__(self, target_shape, **kwargs):\n        super(ReshapeLayer, self).__init__(**kwargs)\n        self.target_shape = target_shape  # e.g., (num_nodes, 3)\n    \n    def call(self, inputs):\n        return tf.reshape(inputs, [-1] + list(self.target_shape))\n    \n    def compute_output_shape(self, input_shape):\n        return (None,) + self.target_shape","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:02:14.906474Z","iopub.status.busy":"2025-05-17T08:02:14.905818Z","iopub.status.idle":"2025-05-17T08:02:14.910428Z","shell.execute_reply":"2025-05-17T08:02:14.909737Z"},"papermill":{"duration":0.01088,"end_time":"2025-05-17T08:02:14.911409","exception":false,"start_time":"2025-05-17T08:02:14.900529","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"a9913668","cell_type":"code","source":"# Custom callback to check for NaN gradients\nclass GradientCheckCallback(tf.keras.callbacks.Callback):\n    def __init__(self, training_data):\n        super(GradientCheckCallback, self).__init__()\n        self.training_data = training_data\n\n    def on_batch_end(self, batch, logs=None):\n        # Fetch the current batch (requires access to the training dataset)\n        # This is tricky because `batch` is just an index, and we need the actual data\n        # For simplicity, this would need a custom training loop or batch access\n        batch_data = next(iter(self.training_data.take(1)))\n        inputs, targets = batch_data\n\n        with tf.GradientTape() as tape:\n            predictions = self.model(inputs, training=True)\n            loss = self.model.compiled_loss(targets, predictions, regularization_losses=self.model.losses)\n\n        grads = tape.gradient(loss, self.model.trainable_weights)\n        for i, g in enumerate(grads):\n            if g is not None and (tf.reduce_any(tf.math.is_nan(g)) or tf.reduce_any(tf.math.is_inf(g))):\n                print(f\"NaN/Inf gradient detected in weight {i}\")","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:02:14.921278Z","iopub.status.busy":"2025-05-17T08:02:14.921058Z","iopub.status.idle":"2025-05-17T08:02:14.926861Z","shell.execute_reply":"2025-05-17T08:02:14.926204Z"},"papermill":{"duration":0.011881,"end_time":"2025-05-17T08:02:14.927989","exception":false,"start_time":"2025-05-17T08:02:14.916108","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"28c65f31","cell_type":"code","source":"# ==============================\n# Data Loading\n# ==============================\n\nDATA_PATH = \"/kaggle/input/stanford-rna-3d-folding/\"\nMSA_DIR = os.path.join(DATA_PATH, \"MSA\")\nMAX_NODES = 1000\nFEATURES = [\n    'res_pos', 'pairing_prob',\n    'freq_A', 'freq_C', 'freq_G', 'freq_U',\n    'resname_A', 'resname_C', 'resname_G', 'resname_U', 'resname_-',\n    'prev_resname_A', 'prev_resname_C', 'prev_resname_G', 'prev_resname_U', 'prev_resname_-',\n    'next_resname_A', 'next_resname_C', 'next_resname_G', 'next_resname_U', 'next_resname_-'\n]\nCOORDINATE_COLUMNS = ['x_1', 'y_1', 'z_1']\nCOORDINATE_COLUMNS_SUBMISSION = [\n    'x_1', 'y_1', 'z_1', 'x_2', 'y_2', 'z_2', 'x_3', 'y_3', 'z_3',\n    'x_4', 'y_4', 'z_4', 'x_5', 'y_5', 'z_5'\n]","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:02:14.938506Z","iopub.status.busy":"2025-05-17T08:02:14.938291Z","iopub.status.idle":"2025-05-17T08:02:14.942353Z","shell.execute_reply":"2025-05-17T08:02:14.941878Z"},"papermill":{"duration":0.010653,"end_time":"2025-05-17T08:02:14.943276","exception":false,"start_time":"2025-05-17T08:02:14.932623","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"82ef5ca8","cell_type":"code","source":"def safe_read_csv(file_path):\n    try:\n        df = pd.read_csv(file_path)\n        print(f\"Successfully loaded {file_path} with shape {df.shape}\")\n        logging.info(f\"Successfully loaded {file_path} with shape {df.shape}\")\n        return df\n    except Exception as e:\n        print(f\"Error: Failed to load {file_path}. Exception: {str(e)}\")\n        logging.error(f\"Failed to load {file_path}. Exception: {str(e)}\")\n        return None\n\ndef safe_read_fasta(file_path):\n    try:\n        return {record.id: str(record.seq) for record in SeqIO.parse(file_path, \"fasta\")}\n    except:\n        return {}\n\ndef load_msa_subset(msa_dir, seq_ids):\n    msa_data = {}\n    for seq_id in seq_ids:\n        msa_file = os.path.join(msa_dir, f\"{seq_id}.MSA.fasta\")\n        if os.path.exists(msa_file):\n            msa_data[seq_id] = safe_read_fasta(msa_file)\n    return msa_data\n\ndef compute_msa_features_for_sequence(msa_dict, target_sequence, max_resid):\n    if not msa_dict or not target_sequence:\n        return pd.DataFrame({\n            'resid': range(1, max_resid + 1),\n            'freq_A': 0, 'freq_C': 0, 'freq_G': 0, 'freq_U': 0\n        })\n    \n    seq_length = min(len(target_sequence), max_resid)\n    msa_sequences = [s[:seq_length] for s in msa_dict.values()]\n    \n    # Convert to TensorFlow tensor and move to GPU\n    msa_array = tf.convert_to_tensor([list(s.ljust(seq_length, '-')) for s in msa_sequences], dtype=tf.string)\n    total_seqs = tf.cast(len(msa_sequences), tf.float32)\n    \n    features = {'resid': np.arange(1, seq_length + 1)}\n    for base in ['A', 'C', 'G', 'U']:\n        freq = tf.reduce_sum(tf.cast(tf.equal(msa_array, base), tf.float32), axis=0) / total_seqs\n        features[f'freq_{base}'] = freq.numpy()\n    \n    return pd.DataFrame(features)\n\n# Top-level function for multiprocessing\ndef process_msa_sequence(args):\n    seq_id, group, msa_df = args\n    msa_df['sequence_id'] = seq_id\n    return msa_df\n\ndef compute_pairing_probabilities(sequence, max_resid):\n    if not isinstance(sequence, str) or not sequence or max_resid <= 0:\n        return np.zeros(max_resid)\n    seq_len = min(len(sequence), max_resid)\n    if seq_len < 2:\n        return np.zeros(max_resid)\n    probs = np.zeros(max_resid)\n    for i in range(seq_len):\n        base = sequence[i]\n        if base in ['G', 'C', 'A', 'U']:\n            probs[i] = 0.1\n    return probs\n\n# Top-level function for multiprocessing\ndef process_secondary_structure(args):\n    seq_id, group, seq_lookup = args\n    sequence = seq_lookup.get(seq_id, '')\n    max_resid = group['resid'].max()\n    probs = compute_pairing_probabilities(sequence, max_resid)\n    return seq_id, probs[:len(group)]\n    \ndef normalize_coordinates(df, coord_cols):\n    for col in coord_cols:\n        stats = df.groupby('sequence_id')[col].agg(['mean', 'std']).reset_index()\n        stats['std'] = stats['std'].replace(0, 1.0)\n        df = df.merge(stats, on='sequence_id')\n        df[col] = (df[col] - df['mean']) / df['std']\n        df[col] = df[col].fillna(0).clip(lower=-10, upper=10)\n        df = df.drop(columns=['mean', 'std'])\n    return df\n\ndef check_data_quality(df, cols):\n    valid_cols = [col for col in cols if col in df.columns]\n    print(f\"Checking columns: {valid_cols}\")\n    print(f\"NaNs in {valid_cols}:\", df[valid_cols].isna().sum())\n    print(f\"Infs in {valid_cols}:\", np.isinf(df[valid_cols]).sum())\n    df[valid_cols] = df[valid_cols].replace([np.inf, -np.inf], 0).fillna(0)\n    return df\n\ndef add_msa_features(df, seq_df, msa_data):\n    seq_lookup = dict(zip(seq_df['ID'], seq_df['sequence']))\n    msa_features = []\n    \n    # Compute all MSA features on GPU in the main process\n    msa_results = []\n    for seq_id, group in df.groupby('sequence_id'):\n        max_resid = group['resid'].max()\n        target_seq = seq_lookup.get(seq_id, '')\n        msa_dict = msa_data.get(seq_id, {})\n        msa_df = compute_msa_features_for_sequence(msa_dict, target_seq, max_resid)\n        msa_results.append((seq_id, group, msa_df))\n    \n    # Parallelize DataFrame operations across CPU cores\n    with ProcessPoolExecutor(max_workers=multiprocessing.cpu_count()) as executor:\n        msa_features = list(executor.map(process_msa_sequence, msa_results))\n    \n    msa_features_df = pd.concat(msa_features)\n    return df.merge(msa_features_df, on=['sequence_id', 'resid'], how='left').fillna(0)\n\ndef add_secondary_structure_features(df, seq_df):\n    seq_lookup = dict(zip(seq_df['ID'], seq_df['sequence']))\n    df['pairing_prob'] = 0.0\n    \n    # Prepare arguments for multiprocessing\n    groups = [(seq_id, group, seq_lookup) for seq_id, group in df.groupby('sequence_id')]\n    \n    # Parallelize across CPU cores\n    with ProcessPoolExecutor(max_workers=multiprocessing.cpu_count()) as executor:\n        results = list(executor.map(process_secondary_structure, groups))\n    \n    # Update DataFrame in one go\n    for seq_id, probs in results:\n        df.loc[df['sequence_id'] == seq_id, 'pairing_prob'] = probs\n    \n    return df\n\ndef add_sequence_context(df):\n    grouped = df.groupby('sequence_id')\n    df['prev_resname'] = grouped['resname'].shift(1).fillna('-')\n    df['next_resname'] = grouped['resname'].shift(-1).fillna('-')\n    return df\n\ndef create_adj_matrix(pairing_probs, threshold=0.05):\n    adj = (pairing_probs[:, None] + pairing_probs[None, :]) > threshold\n    return adj.astype(float)\n\ndef prepare_sequence_data(df, features, max_nodes=MAX_NODES, include_adj=False):\n    node_features = []\n    adj_matrices = []\n    coords = []\n    seq_ids = []\n    res_counts = []\n    for seq_id, group in df.groupby('sequence_id'):\n        feat = group[features].values\n        pairing = group['pairing_prob'].values\n        coord = group[COORDINATE_COLUMNS].values if all(c in df.columns for c in COORDINATE_COLUMNS) else np.zeros((len(group), 3))\n        num_nodes = len(feat)\n        res_counts.append(num_nodes)\n        seq_ids.append(seq_id)\n        if num_nodes > max_nodes:\n            feat = feat[:max_nodes]\n            pairing = pairing[:max_nodes]\n            coord = coord[:max_nodes]\n            num_nodes = max_nodes\n        padded_feat = np.pad(feat, ((0, max_nodes - num_nodes), (0, 0)), mode='constant')\n        padded_coord = np.pad(coord, ((0, max_nodes - num_nodes), (0, 0)), mode='constant')\n        if include_adj:\n            adj = create_adj_matrix(pairing)\n            padded_adj = np.pad(adj, ((0, max_nodes - num_nodes), (0, max_nodes - num_nodes)), mode='constant')\n            adj_matrices.append(padded_adj)\n        node_features.append(padded_feat)\n        coords.append(padded_coord)\n    node_features = np.array(node_features, dtype=np.float32)\n    coords = np.array(coords, dtype=np.float32)\n    # Validate inputs\n    print(f\"node_features shape: {node_features.shape}\")\n    print(f\"coords shape: {coords.shape}\")\n    if np.any(np.isnan(node_features)) or np.any(np.isinf(node_features)):\n        raise ValueError(\"NaNs or Infs detected in node_features\")\n    if np.any(np.isnan(coords)) or np.any(np.isinf(coords)):\n        raise ValueError(\"NaNs or Infs detected in coords\")\n    if include_adj:\n        adj_matrices = np.array(adj_matrices, dtype=np.float32)\n        print(f\"adj_matrices shape: {adj_matrices.shape}\")\n        if np.any(np.isnan(adj_matrices)) or np.any(np.isinf(adj_matrices)):\n            raise ValueError(\"NaNs or Infs detected in adj_matrices\")\n        return node_features, adj_matrices, coords, seq_ids, res_counts\n    return node_features, coords, seq_ids, res_counts\n\n# Build & Model Training with Dense Model (Single GPU)\ndef build_gnn_model(num_features, num_nodes=MAX_NODES):\n    node_input = Input(shape=(num_nodes, num_features), name='node_features', dtype=tf.float32)\n    x = Dense(256, activation='relu')(node_input)\n    x = BatchNormalization()(x)\n    x = Dropout(0.4)(x)\n    x = Dense(128, activation='relu')(x)\n    x = BatchNormalization()(x)\n    x = Dropout(0.4)(x)\n    x = Dense(64, activation='relu')(x)\n    x = BatchNormalization()(x)\n    x = Dropout(0.4)(x)\n    x = tf.keras.layers.Flatten()(x)\n    x = Dense(64, activation='relu')(x)\n    x = BatchNormalization()(x)\n    x = Dropout(0.4)(x)\n    output = Dense(3 * num_nodes)(x)\n    output = ReshapeLayer((num_nodes, 3))(output)\n    return Model(inputs=node_input, outputs=output)\n\n# Update predict_with_dropout to ensure varied predictions\ndef predict_with_dropout(model, X_nodes, num_samples=5, batch_size=4):\n    predictions = []\n    for i in range(num_samples):\n        batch_preds = []\n        for batch_start in range(0, len(X_nodes), batch_size):\n            batch_end = min(batch_start + batch_size, len(X_nodes))\n            batch_X = X_nodes[batch_start:batch_end]\n            tf.random.set_seed(None)\n            np.random.seed(None)\n            pred = model(batch_X, training=True)\n            batch_preds.append(pred.numpy())\n        predictions.append(np.concatenate(batch_preds, axis=0))\n    return np.stack(predictions, axis=1)\n\ndef validate_submission(submission, sample_submission):\n    assert submission.shape == sample_submission.shape, f\"Shape mismatch: {submission.shape} vs {sample_submission.shape}\"\n    assert all(col in submission.columns for col in sample_submission.columns), \"Column mismatch\"\n    assert not submission[COORDINATE_COLUMNS_SUBMISSION].isna().any().any(), \"NaNs in submission\"\n    assert not np.isinf(submission[COORDINATE_COLUMNS_SUBMISSION]).any().any(), \"Infs in submission\"","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:02:14.953789Z","iopub.status.busy":"2025-05-17T08:02:14.953589Z","iopub.status.idle":"2025-05-17T08:02:14.980435Z","shell.execute_reply":"2025-05-17T08:02:14.979907Z"},"papermill":{"duration":0.033301,"end_time":"2025-05-17T08:02:14.981501","exception":false,"start_time":"2025-05-17T08:02:14.948200","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"714b2a2d","cell_type":"code","source":"train_seq = safe_read_csv(os.path.join(DATA_PATH, \"train_sequences.v2.csv\")).rename(columns={'target_id': 'ID'})\ntest_seq = safe_read_csv(os.path.join(DATA_PATH, \"test_sequences.csv\")).rename(columns={'target_id': 'ID'})\ntrain_labels = safe_read_csv(os.path.join(DATA_PATH, \"train_labels.v2.csv\"))\nvalidation_seq = safe_read_csv(os.path.join(DATA_PATH, \"validation_sequences.csv\")).rename(columns={'target_id': 'ID'})\nvalidation_labels = safe_read_csv(os.path.join(DATA_PATH, \"validation_labels.csv\"))\nsample_submission = safe_read_csv(os.path.join(DATA_PATH, \"sample_submission.csv\"))","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:02:14.992138Z","iopub.status.busy":"2025-05-17T08:02:14.991945Z","iopub.status.idle":"2025-05-17T08:02:22.277040Z","shell.execute_reply":"2025-05-17T08:02:22.276332Z"},"papermill":{"duration":7.291638,"end_time":"2025-05-17T08:02:22.278154","exception":false,"start_time":"2025-05-17T08:02:14.986516","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"c83791d7","cell_type":"code","source":"# ==============================\n# Data Preprocessing\n# ==============================\n\ntrain_labels['sequence_id'] = train_labels['ID'].apply(lambda x: x.rsplit('_', 1)[0])\nvalidation_labels['sequence_id'] = validation_labels['ID'].apply(lambda x: x.rsplit('_', 1)[0])\nsample_submission['sequence_id'] = sample_submission['ID'].apply(lambda x: x.rsplit('_', 1)[0])\n\ntrain_data = train_labels.merge(train_seq[['ID', 'sequence']], left_on='sequence_id', right_on='ID', how='left', suffixes=('_label', '_seq'))\nval_data = validation_labels.merge(validation_seq[['ID', 'sequence']], left_on='sequence_id', right_on='ID', how='left', suffixes=('_label', '_seq'))\ntest_data = sample_submission.merge(test_seq[['ID', 'sequence']], left_on='sequence_id', right_on='ID', how='left', suffixes=('_label', '_seq'))\n\ntrain_data = train_data.rename(columns={'ID_label': 'ID_x', 'ID_seq': 'ID_y'})\nval_data = val_data.rename(columns={'ID_label': 'ID_x', 'ID_seq': 'ID_y'})\ntest_data = test_data.rename(columns={'ID_label': 'ID_x', 'ID_seq': 'ID_y'})\n\ndef add_residue_features(df):\n    if 'resid' not in df.columns:\n        df['resid'] = df['ID_x'].str.extract(r'_r(\\d+)$').astype(float)\n    if 'resname' not in df.columns:\n        df['resname'] = df['sequence'].str.split('').str[df['resid'].astype(int)]\n    if 'res_pos' not in df.columns:\n        df['res_pos'] = df.groupby('sequence_id')['resid'].transform(lambda x: x / x.max() if x.max() > 0 else 0)\n    return df\n\ntrain_data = add_residue_features(train_data)\nval_data = add_residue_features(val_data)\ntest_data = add_residue_features(test_data)\n\nfor df in [train_data, val_data, test_data]:\n    df['sequence'] = df['sequence'].fillna('')\n    df['resname'] = df['resname'].fillna('-')\n    df['resid'] = df['resid'].fillna(0)\n    df['res_pos'] = df['res_pos'].fillna(0)\n\nprint(\"train_data columns:\", train_data.columns)\nprint(\"Missing values in train_data:\", train_data.isna().sum())","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:02:22.289224Z","iopub.status.busy":"2025-05-17T08:02:22.289026Z","iopub.status.idle":"2025-05-17T08:02:28.568489Z","shell.execute_reply":"2025-05-17T08:02:28.567630Z"},"papermill":{"duration":6.286798,"end_time":"2025-05-17T08:02:28.569971","exception":false,"start_time":"2025-05-17T08:02:22.283173","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"14c58ac6","cell_type":"code","source":"# ==============================\n# Down-Sample training data\n# ==============================\n\ntrain_data = train_data.sample(n=300000, random_state=42).reset_index(drop=True)","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:02:28.582063Z","iopub.status.busy":"2025-05-17T08:02:28.581830Z","iopub.status.idle":"2025-05-17T08:02:29.186098Z","shell.execute_reply":"2025-05-17T08:02:29.185545Z"},"papermill":{"duration":0.61174,"end_time":"2025-05-17T08:02:29.187473","exception":false,"start_time":"2025-05-17T08:02:28.575733","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"4b73d9f2","cell_type":"code","source":"# ==============================\n# Load and cache MSA features\n# ==============================\n\nseq_ids = set(train_seq['ID']).union(validation_seq['ID'], test_seq['ID'])\nmsa_feature_cache = '/kaggle/working/msa_features.pkl'\nif os.path.exists(msa_feature_cache):\n    with open(msa_feature_cache, 'rb') as f:\n        msa_data = pickle.load(f)\nelse:\n    msa_data = load_msa_subset(MSA_DIR, set(train_data['sequence_id']).union(set(val_data['sequence_id']), set(test_data['sequence_id'])))\n    with open(msa_feature_cache, 'wb') as f:\n        pickle.dump(msa_data, f)","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:02:29.198475Z","iopub.status.busy":"2025-05-17T08:02:29.198247Z","iopub.status.idle":"2025-05-17T08:02:41.622008Z","shell.execute_reply":"2025-05-17T08:02:41.621175Z"},"papermill":{"duration":12.430864,"end_time":"2025-05-17T08:02:41.623638","exception":false,"start_time":"2025-05-17T08:02:29.192774","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"aa1bd201","cell_type":"code","source":"# ==============================\n# Add features\n# ==============================\n\n# Add MSA features\nprint(\"Adding MSA features...\")\ntrain_data = add_msa_features(train_data, train_seq, msa_data)\nval_data = add_msa_features(val_data, validation_seq, msa_data)\ntest_data = add_msa_features(test_data, test_seq, msa_data)\n\n# Add secondary structure features\nprint(\"Adding secondary structure features...\")\ntrain_data = add_secondary_structure_features(train_data, train_seq)\nval_data = add_secondary_structure_features(val_data, validation_seq)\ntest_data = add_secondary_structure_features(test_data, test_seq)\n\n# Add sequence context\nprint(\"Adding sequence context...\")\ntrain_data = add_sequence_context(train_data)\nval_data = add_sequence_context(val_data)\ntest_data = add_sequence_context(test_data)\n\nprint(\"Feature addition completed.\")","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:02:41.635621Z","iopub.status.busy":"2025-05-17T08:02:41.634974Z","iopub.status.idle":"2025-05-17T08:08:31.095203Z","shell.execute_reply":"2025-05-17T08:08:31.094135Z"},"papermill":{"duration":349.471559,"end_time":"2025-05-17T08:08:31.100760","exception":false,"start_time":"2025-05-17T08:02:41.629201","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"83cb7a3a","cell_type":"code","source":"# ==============================\n# Encode categorical features\n# ==============================\n\nresname_categories = ['A', 'C', 'G', 'U', '-']\nfor df in [train_data, val_data, test_data]:\n    df['resname'] = pd.Categorical(df['resname'], categories=resname_categories)\n    df['prev_resname'] = pd.Categorical(df['prev_resname'], categories=resname_categories)\n    df['next_resname'] = pd.Categorical(df['next_resname'], categories=resname_categories)\n    # Create one-hot encoded columns\n    resname_dummies = pd.get_dummies(df['resname'], prefix='resname', dtype=np.uint8)\n    prev_resname_dummies = pd.get_dummies(df['prev_resname'], prefix='prev_resname', dtype=np.uint8)\n    next_resname_dummies = pd.get_dummies(df['next_resname'], prefix='next_resname', dtype=np.uint8)\n    # Add new columns to df\n    df[resname_dummies.columns] = resname_dummies\n    df[prev_resname_dummies.columns] = prev_resname_dummies\n    df[next_resname_dummies.columns] = next_resname_dummies","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:08:31.112353Z","iopub.status.busy":"2025-05-17T08:08:31.111878Z","iopub.status.idle":"2025-05-17T08:08:31.200208Z","shell.execute_reply":"2025-05-17T08:08:31.199689Z"},"papermill":{"duration":0.095388,"end_time":"2025-05-17T08:08:31.201339","exception":false,"start_time":"2025-05-17T08:08:31.105951","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"39381bd8","cell_type":"code","source":"# ==============================\n# Normalize coordinates\n# ==============================\n\ntrain_data = normalize_coordinates(train_data, COORDINATE_COLUMNS)\nval_data = normalize_coordinates(val_data, COORDINATE_COLUMNS)","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:08:31.214530Z","iopub.status.busy":"2025-05-17T08:08:31.213993Z","iopub.status.idle":"2025-05-17T08:08:32.279386Z","shell.execute_reply":"2025-05-17T08:08:32.278552Z"},"papermill":{"duration":1.072763,"end_time":"2025-05-17T08:08:32.280988","exception":false,"start_time":"2025-05-17T08:08:31.208225","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"0b1dbf5c","cell_type":"code","source":"# ==============================\n# Check Data Quality\n# ==============================\n\ntrain_data = check_data_quality(train_data, FEATURES + COORDINATE_COLUMNS)\nval_data = check_data_quality(val_data, FEATURES + COORDINATE_COLUMNS)\ntest_data = check_data_quality(test_data, FEATURES)\n\nprint(\"train_data columns after feature engineering:\", train_data.columns)\nprint(\"Missing values in train_data:\", train_data.isna().sum())","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:08:32.292527Z","iopub.status.busy":"2025-05-17T08:08:32.292251Z","iopub.status.idle":"2025-05-17T08:08:32.565792Z","shell.execute_reply":"2025-05-17T08:08:32.564623Z"},"papermill":{"duration":0.280682,"end_time":"2025-05-17T08:08:32.567019","exception":false,"start_time":"2025-05-17T08:08:32.286337","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"010cf8bb","cell_type":"code","source":"# ==============================\n# Prepare GNN data\n# ==============================\n\nX_train_nodes, y_train, train_seq_ids, train_res_counts = prepare_sequence_data(train_data, FEATURES, include_adj=False)\nX_val_nodes, y_val, val_seq_ids, val_res_counts = prepare_sequence_data(val_data, FEATURES, include_adj=False)\nX_test_nodes, _, test_seq_ids, test_res_counts = prepare_sequence_data(test_data, FEATURES, include_adj=False)","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:08:32.579513Z","iopub.status.busy":"2025-05-17T08:08:32.579255Z","iopub.status.idle":"2025-05-17T08:08:40.946147Z","shell.execute_reply":"2025-05-17T08:08:40.945196Z"},"papermill":{"duration":8.374485,"end_time":"2025-05-17T08:08:40.947577","exception":false,"start_time":"2025-05-17T08:08:32.573092","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"6c0084ee","cell_type":"code","source":"# Debug NaNs before training\nprint(\"NaNs in X_train_nodes:\", np.isnan(X_train_nodes).sum())\n\nprint(\"NaNs in y_train:\", np.isnan(y_train).sum())","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:08:40.960337Z","iopub.status.busy":"2025-05-17T08:08:40.959916Z","iopub.status.idle":"2025-05-17T08:08:41.133126Z","shell.execute_reply":"2025-05-17T08:08:41.132329Z"},"papermill":{"duration":0.180345,"end_time":"2025-05-17T08:08:41.134310","exception":false,"start_time":"2025-05-17T08:08:40.953965","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"47446ef4","cell_type":"code","source":"print(f\"X_train_nodes min/max: {X_train_nodes.min()}, {X_train_nodes.max()}\")\nprint(f\"y_train min/max: {y_train.min()}, {y_train.max()}\")\nprint(f\"NaNs in X_train_nodes: {np.isnan(X_train_nodes).sum()}\")\nprint(f\"NaNs in y_train: {np.isnan(y_train).sum()}\")","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:08:41.146671Z","iopub.status.busy":"2025-05-17T08:08:41.146414Z","iopub.status.idle":"2025-05-17T08:08:41.450213Z","shell.execute_reply":"2025-05-17T08:08:41.449395Z"},"papermill":{"duration":0.311432,"end_time":"2025-05-17T08:08:41.451626","exception":false,"start_time":"2025-05-17T08:08:41.140194","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"1c847a31","cell_type":"code","source":"# Normalize Inputs and Targets\nscaler_nodes = StandardScaler()\nX_train_nodes_reshaped = X_train_nodes.reshape(-1, X_train_nodes.shape[-1])\nX_val_nodes_reshaped = X_val_nodes.reshape(-1, X_val_nodes.shape[-1])\nX_test_nodes_reshaped = X_test_nodes.reshape(-1, X_test_nodes.shape[-1])\n\nX_train_nodes_scaled = scaler_nodes.fit_transform(X_train_nodes_reshaped).reshape(X_train_nodes.shape)\nX_val_nodes_scaled = scaler_nodes.transform(X_val_nodes_reshaped).reshape(X_val_nodes.shape)\nX_test_nodes_scaled = scaler_nodes.transform(X_test_nodes_reshaped).reshape(X_test_nodes.shape)\n\nscaler_targets = StandardScaler()\ny_train_reshaped = y_train.reshape(-1, y_train.shape[-1])\ny_val_reshaped = y_val.reshape(-1, y_val.shape[-1])\n\ny_train_scaled = scaler_targets.fit_transform(y_train_reshaped).reshape(y_train.shape)\ny_val_scaled = scaler_targets.transform(y_val_reshaped).reshape(y_val.shape)","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:08:41.463686Z","iopub.status.busy":"2025-05-17T08:08:41.463473Z","iopub.status.idle":"2025-05-17T08:08:43.459129Z","shell.execute_reply":"2025-05-17T08:08:43.458300Z"},"papermill":{"duration":2.00351,"end_time":"2025-05-17T08:08:43.460754","exception":false,"start_time":"2025-05-17T08:08:41.457244","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"255b64f2","cell_type":"code","source":"# Debug normalized data\nprint(f\"X_train_nodes_scaled min/max: {X_train_nodes_scaled.min()}, {X_train_nodes_scaled.max()}\")\nprint(f\"y_train_scaled min/max: {y_train_scaled.min()}, {y_train_scaled.max()}\")\nprint(f\"NaNs in X_train_nodes_scaled: {np.isnan(X_train_nodes_scaled).sum()}\")\nprint(f\"NaNs in y_train_scaled: {np.isnan(y_train_scaled).sum()}\")","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:08:43.474903Z","iopub.status.busy":"2025-05-17T08:08:43.474476Z","iopub.status.idle":"2025-05-17T08:08:43.776788Z","shell.execute_reply":"2025-05-17T08:08:43.775906Z"},"papermill":{"duration":0.309752,"end_time":"2025-05-17T08:08:43.777971","exception":false,"start_time":"2025-05-17T08:08:43.468219","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"259084f2","cell_type":"code","source":"# Convert training and validation data to tf.data.Dataset\ntrain_dataset = tf.data.Dataset.from_tensor_slices((X_train_nodes_scaled, y_train_scaled)).batch(16).prefetch(tf.data.AUTOTUNE)\nval_dataset = tf.data.Dataset.from_tensor_slices((X_val_nodes_scaled, y_val_scaled)).batch(16).prefetch(tf.data.AUTOTUNE)","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:08:43.790738Z","iopub.status.busy":"2025-05-17T08:08:43.790069Z","iopub.status.idle":"2025-05-17T08:08:44.635553Z","shell.execute_reply":"2025-05-17T08:08:44.634896Z"},"papermill":{"duration":0.852936,"end_time":"2025-05-17T08:08:44.636849","exception":false,"start_time":"2025-05-17T08:08:43.783913","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"3ad1d462","cell_type":"code","source":"# ==============================\n# Build & Model Training\n# ==============================\n\nmodel = build_gnn_model(num_features=len(FEATURES))\nmodel.compile(optimizer=AdamW(learning_rate=1e-4, weight_decay=1e-4, clipnorm=1.0), \n              loss=lambda y_true, y_pred: tf.keras.losses.mse(y_true, y_pred) + 1e-6)\n\nlr_scheduler = ReduceLROnPlateau(patience=5, factor=0.8)\n\n# Check initial predictions for NaNs\ntest_batch_nodes = X_train_nodes_scaled[:8]\ninitial_pred = model.predict(test_batch_nodes, batch_size=8)\nprint(f\"Initial predictions min/max: {initial_pred.min()}, {initial_pred.max()}\")\nprint(f\"Initial predictions NaNs: {np.isnan(initial_pred).sum()}\")\n\n# Train the model using tf.data.Dataset\nmodel.fit(\n    train_dataset,\n    validation_data=val_dataset,\n    epochs=10,\n    callbacks=[lr_scheduler, GradientCheckCallback(train_dataset)],\n    verbose=2\n)","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:08:44.649925Z","iopub.status.busy":"2025-05-17T08:08:44.649078Z","iopub.status.idle":"2025-05-17T08:30:21.899665Z","shell.execute_reply":"2025-05-17T08:30:21.898864Z"},"papermill":{"duration":1297.261679,"end_time":"2025-05-17T08:30:21.904399","exception":false,"start_time":"2025-05-17T08:08:44.642720","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"d9867f25","cell_type":"code","source":"# ==============================\n# Generate Predictions\n# ==============================\n\ny_test_pred_scaled = predict_with_dropout(model, X_test_nodes_scaled, num_samples=5, batch_size=4)\n\n# Inverse transform predictions to original scale\ny_test_pred_reshaped = y_test_pred_scaled.reshape(-1, y_test_pred_scaled.shape[-1])\ny_test_pred = scaler_targets.inverse_transform(y_test_pred_reshaped).reshape(y_test_pred_scaled.shape)\n\n# Handle NaNs/Infs and clip coordinates\ny_test_pred = np.nan_to_num(y_test_pred, nan=0.0, posinf=0.0, neginf=0.0)\ny_test_pred = np.clip(y_test_pred, -100, 100)\n\nprint(f\"y_test_pred shape: {y_test_pred.shape}\")\nprint(f\"y_test_pred min/max: {np.nanmin(y_test_pred)}, {np.nanmax(y_test_pred)}\")\nprint(f\"y_test_pred NaNs: {np.isnan(y_test_pred).sum()}\")\nprint(f\"y_test_pred[0, :, 0, :]: {y_test_pred[0, :, 0, :]}\")","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:30:21.919364Z","iopub.status.busy":"2025-05-17T08:30:21.918738Z","iopub.status.idle":"2025-05-17T08:30:22.304496Z","shell.execute_reply":"2025-05-17T08:30:22.303769Z"},"papermill":{"duration":0.39429,"end_time":"2025-05-17T08:30:22.305622","exception":false,"start_time":"2025-05-17T08:30:21.911332","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"1d1443f0","cell_type":"code","source":"# Clean up\ngc.collect()\ntf.keras.backend.clear_session()","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:30:22.511200Z","iopub.status.busy":"2025-05-17T08:30:22.510614Z","iopub.status.idle":"2025-05-17T08:30:22.799340Z","shell.execute_reply":"2025-05-17T08:30:22.798604Z"},"papermill":{"duration":0.297866,"end_time":"2025-05-17T08:30:22.800631","exception":false,"start_time":"2025-05-17T08:30:22.502765","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"1e391e35","cell_type":"code","source":"# ==============================\n# Prepare Submission\n# ==============================\n\nsubmission = sample_submission.copy()\npred_dict = {}\n\nfor seq_idx, seq_id in enumerate(test_seq_ids):\n    num_residues = test_res_counts[seq_idx]\n    print(f\"Test sequence {seq_id}: num_residues={num_residues}\")  # Debug hidden test set\n    coords = y_test_pred[seq_idx, :, :min(num_residues, MAX_NODES), :]\n    print(f\"Sequence {seq_id}: num_residues={num_residues}, coords shape={coords.shape}\")\n    for sample_idx in range(coords.shape[0]):\n        for res_idx in range(num_residues):\n            res_id = f\"{seq_id}_{res_idx + 1}\"\n            if res_idx < coords.shape[1]:\n                pred_dict[(res_id, sample_idx)] = coords[sample_idx, res_idx, :]\n            else:\n                mean_coords = np.mean(coords[sample_idx, :num_residues, :], axis=0)\n                pred_dict[(res_id, sample_idx)] = mean_coords\n\n# Assign predictions to submission\nfor sample_idx in range(5):\n    for coord_idx, coord_name in enumerate(['x', 'y', 'z']):\n        col = f\"{coord_name}_{sample_idx + 1}\"\n        submission[col] = submission['ID'].apply(\n            lambda id: pred_dict.get((id, sample_idx), [0, 0, 0])[coord_idx]\n        )","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:30:22.320304Z","iopub.status.busy":"2025-05-17T08:30:22.319917Z","iopub.status.idle":"2025-05-17T08:30:22.366767Z","shell.execute_reply":"2025-05-17T08:30:22.366129Z"},"papermill":{"duration":0.055275,"end_time":"2025-05-17T08:30:22.367891","exception":false,"start_time":"2025-05-17T08:30:22.312616","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"fef10cfa-5e39-49a8-b02d-26207e5c5ba2","cell_type":"code","source":"# Ensure no NaNs or Infs\nsubmission[COORDINATE_COLUMNS_SUBMISSION] = submission[COORDINATE_COLUMNS_SUBMISSION].fillna(0)\nsubmission[COORDINATE_COLUMNS_SUBMISSION] = np.nan_to_num(\n    submission[COORDINATE_COLUMNS_SUBMISSION].values, nan=0.0, posinf=0.0, neginf=0.0\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"aff22212-38c9-4028-8a7b-f91c1e5e5454","cell_type":"code","source":"# Debug submission coordinates\nprint(f\"Submission coordinates min/max: {submission[COORDINATE_COLUMNS_SUBMISSION].min().min()}, {submission[COORDINATE_COLUMNS_SUBMISSION].max().max()}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"f715a875","cell_type":"code","source":"# ==============================\n# Validate and save submission\n# ==============================\n\nvalidate_submission(submission, sample_submission)\nsubmission = submission[['ID', 'resname', 'resid'] + COORDINATE_COLUMNS_SUBMISSION]\nsubmission.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:30:22.382541Z","iopub.status.busy":"2025-05-17T08:30:22.382108Z","iopub.status.idle":"2025-05-17T08:30:22.427951Z","shell.execute_reply":"2025-05-17T08:30:22.427223Z"},"papermill":{"duration":0.054163,"end_time":"2025-05-17T08:30:22.429232","exception":false,"start_time":"2025-05-17T08:30:22.375069","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"a257b363","cell_type":"code","source":"submission.head()","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:30:22.445376Z","iopub.status.busy":"2025-05-17T08:30:22.444809Z","iopub.status.idle":"2025-05-17T08:30:22.471169Z","shell.execute_reply":"2025-05-17T08:30:22.470556Z"},"papermill":{"duration":0.036001,"end_time":"2025-05-17T08:30:22.472273","exception":false,"start_time":"2025-05-17T08:30:22.436272","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"e2dce571","cell_type":"code","source":"print(f\"Submission shape: {submission.shape}\")\nprint(f\"Submission NaNs: {submission[COORDINATE_COLUMNS_SUBMISSION].isna().sum()}\")","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:30:22.487541Z","iopub.status.busy":"2025-05-17T08:30:22.487274Z","iopub.status.idle":"2025-05-17T08:30:22.494117Z","shell.execute_reply":"2025-05-17T08:30:22.493485Z"},"papermill":{"duration":0.015575,"end_time":"2025-05-17T08:30:22.495217","exception":false,"start_time":"2025-05-17T08:30:22.479642","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"36d8ef1f","cell_type":"code","source":"# ==============================\n# Compute TM Score\n# ==============================\n\n# Function to compute d0 based on L_ref\ndef compute_d0(L_ref):\n    if L_ref < 12:\n        return 0.3\n    elif 12 <= L_ref <= 15:\n        return 0.4\n    elif 16 <= L_ref <= 19:\n        return 0.5\n    elif 20 <= L_ref <= 23:\n        return 0.6\n    elif 24 <= L_ref <= 29:\n        return 0.7\n    else:\n        return 1.24 * (L_ref - 15) ** (1/3) - 1.8\n\n# Kabsch algorithm for rigid-body alignment (simplified)\ndef kabsch_align(pred_coords, ref_coords):\n    # Center the coordinates\n    pred_centroid = np.mean(pred_coords, axis=0)\n    ref_centroid = np.mean(ref_coords, axis=0)\n    pred_centered = pred_coords - pred_centroid\n    ref_centered = ref_coords - ref_centroid\n    \n    # Compute the covariance matrix\n    H = np.dot(pred_centered.T, ref_centered)\n    \n    # Singular Value Decomposition\n    U, _, Vt = np.linalg.svd(H)\n    \n    # Rotation matrix\n    R = np.dot(Vt.T, U.T)\n    \n    # Ensure a right-handed coordinate system\n    if np.linalg.det(R) < 0:\n        Vt[-1, :] *= -1\n        R = np.dot(Vt.T, U.T)\n    \n    # Translate and rotate predicted coordinates\n    pred_aligned = np.dot(pred_centered, R) + ref_centroid\n    return pred_aligned\n\n# Function to compute TM-score\ndef compute_tm_score(pred_coords, ref_coords):\n    # Remove rows with NaN in ref_coords\n    mask = ~np.isnan(ref_coords).any(axis=1)\n    pred_coords = pred_coords[mask]\n    ref_coords = ref_coords[mask]\n    \n    L_ref = len(ref_coords)\n    if L_ref == 0:\n        return 0.0\n    \n    # Align the structures\n    pred_aligned = kabsch_align(pred_coords, ref_coords)\n    \n    # Compute distances between aligned residues\n    distances = np.sqrt(np.sum((pred_aligned - ref_coords) ** 2, axis=1))\n    \n    # Compute d0\n    d0 = compute_d0(L_ref)\n    \n    # Compute TM-score\n    tm_score = np.sum(1 / (1 + (distances / d0) ** 2)) / L_ref\n    return tm_score","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:30:22.815837Z","iopub.status.busy":"2025-05-17T08:30:22.815608Z","iopub.status.idle":"2025-05-17T08:30:22.823300Z","shell.execute_reply":"2025-05-17T08:30:22.822554Z"},"papermill":{"duration":0.016644,"end_time":"2025-05-17T08:30:22.824468","exception":false,"start_time":"2025-05-17T08:30:22.807824","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"7020f360","cell_type":"code","source":"# Evaluate TM-score using validation set\n# Map validation sequences to submission predictions\nsubmission['sequence_id'] = submission['ID'].apply(lambda x: x.split('_')[0])\n\n# Map validation sequences to submission predictions\nval_submission = submission[submission['sequence_id'].isin(val_data['sequence_id'].unique())].copy()\ntm_scores = []\n\n# For each sequence in the validation set\nfor seq_id in val_data['sequence_id'].unique():\n    # Get reference coordinates from validation_labels\n    ref_df = val_data[val_data['sequence_id'] == seq_id][['x_1', 'y_1', 'z_1']].values\n    \n    # Get predicted coordinates (5 structures) from submission\n    pred_df = val_submission[val_submission['sequence_id'] == seq_id]\n    \n    # Compute TM-score for each of the 5 predictions\n    seq_tm_scores = []\n    for sample_idx in range(5):\n        pred_coords = pred_df[[f'x_{sample_idx+1}', f'y_{sample_idx+1}', f'z_{sample_idx+1}']].values\n        tm_score = compute_tm_score(pred_coords, ref_df)\n        seq_tm_scores.append(tm_score)\n    \n    # Take the best TM-score for this sequence\n    best_tm_score = max(seq_tm_scores)\n    tm_scores.append(best_tm_score)\n    print(f\"Sequence {seq_id}: Best TM-score = {best_tm_score:.4f}\")\n\n# Compute the average TM-score across all sequences\naverage_tm_score = np.mean(tm_scores)\nprint(f\"Average TM-score across validation sequences: {average_tm_score:.4f}\")","metadata":{"execution":{"iopub.execute_input":"2025-05-17T08:30:22.839382Z","iopub.status.busy":"2025-05-17T08:30:22.838993Z","iopub.status.idle":"2025-05-17T08:30:23.019388Z","shell.execute_reply":"2025-05-17T08:30:23.018623Z"},"papermill":{"duration":0.189376,"end_time":"2025-05-17T08:30:23.020691","exception":false,"start_time":"2025-05-17T08:30:22.831315","status":"completed"},"tags":[]},"outputs":[],"execution_count":null}]}