{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":87793,"databundleVersionId":12276181,"sourceType":"competition"}],"dockerImageVersionId":30918,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport os","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-29T03:23:11.685657Z","iopub.execute_input":"2025-05-29T03:23:11.686051Z","iopub.status.idle":"2025-05-29T03:23:12.108763Z","shell.execute_reply.started":"2025-05-29T03:23:11.686017Z","shell.execute_reply":"2025-05-29T03:23:12.107471Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"file_paths = {\n    'train_sequences_v1': '/kaggle/input/stanford-rna-3d-folding/train_sequences.csv',\n    'train_sequences_v2': '/kaggle/input/stanford-rna-3d-folding/train_sequences.v2.csv',\n    'train_labels_v1': '/kaggle/input/stanford-rna-3d-folding/train_labels.csv',\n    'train_labels_v2': '/kaggle/input/stanford-rna-3d-folding/train_labels.v2.csv', \n    'validation_sequences': '/kaggle/input/stanford-rna-3d-folding/validation_sequences.csv',\n    'validation_labels': '/kaggle/input/stanford-rna-3d-folding/validation_labels.csv',\n    'test_sequences': '/kaggle/input/stanford-rna-3d-folding/test_sequences.csv',\n    'msa_folder': '/kaggle/input/stanford-rna-3d-folding/MSA/'\n}\n\nfor path in file_paths.keys():\n    if os.path.exists(os.path.join(os.getcwd(), file_paths[path])):\n        train_sequences = pd.concat([pd.read_csv(file_paths['train_sequences_v1']), pd.read_csv(file_paths['train_sequences_v2'])], axis=0)\n        train_labels = pd.concat([pd.read_csv(file_paths['train_labels_v1']), pd.read_csv(file_paths['train_labels_v2'])], axis=0)\n        validation_sequences = pd.read_csv(file_paths['validation_sequences'])\n        validation_labels = pd.read_csv(file_paths['validation_labels'])\n        test_sequences = pd.read_csv(file_paths['test_sequences'])\n        MSA_FOLDER = file_paths['msa_folder']\n    else:\n        pass\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T03:23:12.110623Z","iopub.execute_input":"2025-05-29T03:23:12.111263Z","iopub.status.idle":"2025-05-29T03:23:59.139431Z","shell.execute_reply.started":"2025-05-29T03:23:12.111216Z","shell.execute_reply":"2025-05-29T03:23:59.138114Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MAX_SEQ_LENGTH = 10","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T03:23:59.141611Z","iopub.execute_input":"2025-05-29T03:23:59.141999Z","iopub.status.idle":"2025-05-29T03:23:59.146936Z","shell.execute_reply.started":"2025-05-29T03:23:59.141963Z","shell.execute_reply":"2025-05-29T03:23:59.145271Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"max sequence: \", MAX_SEQ_LENGTH)\n\nclass FastaAlignment:  \n    class SeqRecord:\n        def __init__(self, seq, id=\"\", description=\"\"):\n            self.seq = seq\n            self.id = id\n            self.description = description\n    \n    def __init__(self):\n        self.records = []\n    \n    def append(self, record):\n        self.records.append(record)\n    \n    def get_alignment_length(self):\n        if not self.records:\n            return 0\n        return len(self.records[0].seq)\n    \n    def __len__(self):\n        return len(self.records)\n    \n    def __getitem__(self, index):\n        return self.records[index]\n\ndef read_fasta(file_path):\n    alignment = FastaAlignment()\n    \n    try:\n        with open(file_path, 'r') as fasta_file:\n            current_id = \"\"\n            current_description = \"\"\n            current_seq = \"\"\n            \n            for line in fasta_file:\n                line = line.strip()\n                \n                if not line: \n                    continue\n                    \n                if line.startswith('>'):  \n                    if current_seq:\n                        record = FastaAlignment.SeqRecord(current_seq, current_id, current_description)\n                        alignment.append(record)\n                    \n                    header_parts = line[1:].split(maxsplit=1)\n                    current_id = header_parts[0]\n                    current_description = header_parts[1] if len(header_parts) > 1 else \"\"\n                    current_seq = \"\"\n                else:  \n                    current_seq += line\n            \n            if current_seq:\n                record = FastaAlignment.SeqRecord(current_seq, current_id, current_description)\n                alignment.append(record)\n    except:\n        return None\n    \n    if len(alignment) > 0:\n        first_len = len(alignment[0].seq)\n        for record in alignment.records:\n            if len(record.seq) != first_len:\n                print(f\"Warning: Sequences in {file_path} have different lengths.\")\n    \n    return alignment if len(alignment) > 0 else None\n\nclass RNAData:\n    def __init__(self, input_sequences, input_labels):\n        try:\n            self.sequence = [row[-1]['sequence'] for row in input_sequences.iterrows()]\n            self.target_id = [row[-1]['target_id'] for row in input_sequences.iterrows()]\n        except:\n            self.sequence = []\n            self.target_id = []\n        \n        self.coords_dict = {}\n        \n        if input_labels is not None:\n            try:\n                for target_id in self.target_id:\n                    matching_rows = input_labels[input_labels['ID'].str.startswith(target_id)]\n                    if not matching_rows.empty:\n                        coords = matching_rows[['x_1', 'y_1', 'z_1']].values\n                        self.coords_dict[target_id] = coords\n            except:\n                pass\n\n    def one_hot_encode(self, sequence):\n        nucleotide_map = {'A': [1, 0, 0, 0], 'C': [0, 1, 0, 0], 'G': [0, 0, 1, 0], 'U': [0, 0, 0, 1]}\n        \n        valid_nucleotides = [nucleotide for nucleotide in sequence if nucleotide in nucleotide_map]\n        if not valid_nucleotides:\n            return np.zeros((1, 4))\n        encoding = [nucleotide_map[nucleotide] for nucleotide in sequence if nucleotide in nucleotide_map]\n        return np.array(encoding)\n\n    def filter_extreme_coordinates(self):\n        removed_targets = []\n        for target_id in list(self.coords_dict.keys()):\n            coords = self.coords_dict[target_id]\n            if np.max(np.abs(coords)) > 1e10:\n                print(f\"Filtering out target {target_id} with extreme coordinate values\")\n                removed_targets.append(target_id)\n                del self.coords_dict[target_id]\n        \n        if removed_targets:\n            indices_to_keep = [i for i, tid in enumerate(self.target_id) if tid not in removed_targets]\n            self.target_id = [self.target_id[i] for i in indices_to_keep]\n            self.sequence = [self.sequence[i] for i in indices_to_keep]\n            print(f\"Removed {len(removed_targets)} targets with extreme coordinate values\")\n        \n        return removed_targets\n\n    def truncate_sequence(self, sequence, coords, max_length):\n        if len(sequence) <= max_length:\n            return sequence, coords\n        \n        truncated_seq = sequence[:max_length]\n        \n        if coords is not None and len(coords) > 0:\n            truncated_coords = coords[:max_length] if len(coords) > max_length else coords\n        else:\n            truncated_coords = coords\n            \n        return truncated_seq, truncated_coords\n\nclass Geodata(RNAData):\n    def __init__(self, input_sequences, input_labels):\n        super(Geodata, self).__init__(input_sequences, input_labels)\n        if input_labels is not None:\n            try:\n                coords_array = input_labels[['x_1', 'y_1', 'z_1']].to_numpy()\n                valid_mask = np.isfinite(coords_array).all(axis=1)\n                coords_array = coords_array[valid_mask]\n                finite_mask = np.all(np.abs(coords_array) <= 1e10, axis=1)\n                coords_array = coords_array[finite_mask]\n                \n                mean = np.mean(coords_array, axis=0)\n                std = np.std(coords_array, axis=0)\n\n                self.normalization_params = {'mean': mean, 'std': std}\n            except:\n                self.normalization_params = {'mean': np.zeros(3), 'std': np.ones(3)}\n        else:\n            self.normalization_params = {'mean': np.zeros(3), 'std': np.ones(3)}\n\n    def predict_secondary_structure_simple(self, sequence):\n        seq_length = len(sequence)\n        dot_bracket = '.' * seq_length\n        pairing_matrix = np.zeros((seq_length, seq_length))\n        bpp_matrix = np.zeros((seq_length, seq_length))\n        \n        complement = {'A': 'U', 'U': 'A', 'G': 'C', 'C': 'G'}\n        \n        pairs = []\n        for i in range(seq_length - 4): \n            for j in range(i + 4, seq_length):\n                if sequence[i] in complement and sequence[j] == complement[sequence[i]]:\n                    if sequence[i] in ['G', 'C']:\n                        score = 0.8\n                    else:\n                        score = 0.6  \n                    \n                    distance_factor = max(0.1, 1.0 - (j - i) / seq_length)\n                    final_score = score * distance_factor\n                    \n                    pairs.append((i, j, final_score))\n        \n        pairs.sort(key=lambda x: x[2], reverse=True)\n        used_positions = set()\n        \n        dot_bracket_list = list(dot_bracket)\n        for i, j, score in pairs:\n            if i not in used_positions and j not in used_positions:\n                dot_bracket_list[i] = '('\n                dot_bracket_list[j] = ')'\n                pairing_matrix[i, j] = 1\n                pairing_matrix[j, i] = 1\n                bpp_matrix[i, j] = score\n                bpp_matrix[j, i] = score\n                used_positions.add(i)\n                used_positions.add(j)\n        \n        dot_bracket = ''.join(dot_bracket_list)\n        \n        return {\n            'dot_bracket': dot_bracket,\n            'mfe': -len([c for c in dot_bracket if c in '()']) * 2.0,  \n            'pairing_matrix': pairing_matrix,\n            'bpp_matrix': bpp_matrix\n        }\n    \n    def calculate_position_specific_features(self, sequence, sec_struct_data):\n        seq_length = len(sequence)\n        features = np.zeros((seq_length, 10))\n        for i in range(seq_length):\n            window_start = max(0, i-2)\n            window_end = min(seq_length, i+3)\n            window = sequence[window_start:window_end]\n            \n            features[i, 0] = window.count('A') / len(window)\n            features[i, 1] = window.count('C') / len(window)\n            features[i, 2] = window.count('G') / len(window)\n            features[i, 3] = window.count('U') / len(window)\n            \n            features[i, 4] = 1 if sec_struct_data['dot_bracket'][i] == '(' else 0\n            features[i, 5] = 1 if sec_struct_data['dot_bracket'][i] == ')' else 0\n            features[i, 6] = 1 if sec_struct_data['dot_bracket'][i] == '.' else 0\n            \n            features[i, 7] = np.sum(sec_struct_data['bpp_matrix'][i]) if hasattr(sec_struct_data['bpp_matrix'], 'shape') else 0\n            features[i, 8] = i / seq_length\n            features[i, 9] = (seq_length - i) / seq_length\n        \n        return features\n        \n    def augment_geo_data(self, coords):\n        if coords.shape[0] == 0:\n            return coords\n            \n        if coords.ndim == 1:\n            coords = coords.reshape(1, -1)\n        \n        if np.max(np.abs(coords)) > 1e10:\n            return coords\n        \n        try:\n            theta_x = np.random.uniform(0, 2*np.pi)\n            theta_y = np.random.uniform(0, 2*np.pi)\n            theta_z = np.random.uniform(0, 2*np.pi)\n            Rx = np.array([\n                [1, 0, 0],\n                [0, np.cos(theta_x), -np.sin(theta_x)],\n                [0, np.sin(theta_x), np.cos(theta_x)]\n            ])\n            Ry = np.array([\n                [np.cos(theta_y), 0, np.sin(theta_y)],\n                [0, 1, 0],\n                [-np.sin(theta_y), 0, np.cos(theta_y)]\n            ])\n                        \n            Rz = np.array([\n                [np.cos(theta_z), -np.sin(theta_z), 0],\n                [np.sin(theta_z), np.cos(theta_z), 0],\n                [0, 0, 1]\n            ])\n            \n            R = np.dot(Rz, np.dot(Ry, Rx))\n            rotated_coords = coords.copy()\n            \n            if coords.shape[1] >= 3:\n                valid_mask = ~np.all(coords == 0, axis=1)\n                \n                if np.any(valid_mask):\n                    center = np.mean(coords[valid_mask], axis=0)\n                    centered_coords = coords[valid_mask] - center\n                    rotated_points = np.dot(centered_coords, R.T)\n                    rotated_coords[valid_mask] = rotated_points + center\n            \n            return rotated_coords\n        except:\n            return coords\n\n    def normalize_coordinates(self, coords):\n        if coords.shape[0] == 0:\n            return coords\n        \n        if coords.ndim == 1:\n            coords = np.expand_dims(coords, axis=0)\n        \n        try:\n            normalized_coords = (coords - self.normalization_params['mean']) / self.normalization_params['std']\n            return normalized_coords\n        except:\n            return coords\n            \n    def prepare_geo_data(self, augment=False, num_augmentations=0, MAX_SEQ_LENGTH=None):\n        if MAX_SEQ_LENGTH is None:\n            MAX_SEQ_LENGTH = globals().get('MAX_SEQ_LENGTH', 500)\n        \n        self.filter_extreme_coordinates()\n        X_geo = []\n        Y_coords = []\n        metadata = []\n        \n        for target_id, seq in tqdm(zip(self.target_id, self.sequence), desc=\"Processing geometric data\", total=len(self.sequence)):\n            coords = self.coords_dict.get(target_id, None)\n            \n            truncated_seq, truncated_coords = self.truncate_sequence(seq, coords, MAX_SEQ_LENGTH)\n            \n            predicted_structure = self.predict_secondary_structure_simple(truncated_seq)\n            position_features = self.calculate_position_specific_features(truncated_seq, predicted_structure)\n            \n            if truncated_coords is not None:\n                normalized_coords = self.normalize_coordinates(truncated_coords)\n                coords_pad = np.zeros((MAX_SEQ_LENGTH, 3))\n                valid_len = min(len(normalized_coords), MAX_SEQ_LENGTH)\n                coords_pad[:valid_len] = normalized_coords[:valid_len]\n                Y_coords.append(coords_pad)\n            else:\n                coords_pad = np.zeros((MAX_SEQ_LENGTH, 3))\n                Y_coords.append(coords_pad)\n            \n            one_hot = self.one_hot_encode(truncated_seq)\n            min_length = min(len(one_hot), len(position_features))\n            one_hot = one_hot[:min_length]\n            position_features = position_features[:min_length]\n            \n            if position_features.shape[0] > 0 and one_hot.shape[0] > 0:\n                combined_features = np.concatenate((one_hot, position_features), axis=1)\n                padded_features = np.zeros((MAX_SEQ_LENGTH, combined_features.shape[1]))\n                padded_features[:len(combined_features)] = combined_features\n                X_geo.append(padded_features)\n                \n                metadata.append({\n                    'target_id': target_id,\n                    'seq_length': len(truncated_seq),\n                    'original_seq_length': len(seq),\n                    'truncated': len(seq) > MAX_SEQ_LENGTH\n                })\n                \n                if augment and truncated_coords is not None:\n                    if len(truncated_coords) > 0 and np.max(np.abs(truncated_coords)) <= 1e10:\n                        for aug_idx in range(num_augmentations):\n                            augmented_coords = self.augment_geo_data(truncated_coords)\n                            normalized_aug_coords = self.normalize_coordinates(augmented_coords)\n                            \n                            augmented_pad = np.zeros((MAX_SEQ_LENGTH, 3))\n                            valid_len = min(len(normalized_aug_coords), MAX_SEQ_LENGTH)\n                            augmented_pad[:valid_len] = normalized_aug_coords[:valid_len]\n                            \n                            X_geo.append(padded_features)  \n                            Y_coords.append(augmented_pad)\n                            \n                            metadata.append({\n                                'target_id': target_id,\n                                'seq_length': len(truncated_seq),\n                                'original_seq_length': len(seq),\n                                'truncated': len(seq) > MAX_SEQ_LENGTH,\n                                'augmentation': aug_idx + 1\n                            })\n        \n        X_geo_array = np.array(X_geo)\n        Y_coords_array = np.array(Y_coords)\n        \n        if np.isnan(X_geo_array).any():\n            print(f\"Found {np.isnan(X_geo_array).sum()} NaN values in X_geo. Replacing with zeros.\")\n            X_geo_array = np.nan_to_num(X_geo_array)\n            \n        if np.isnan(Y_coords_array).any():\n            print(f\"Found {np.isnan(Y_coords_array).sum()} NaN values in Y_coords. Replacing with zeros.\")\n            Y_coords_array = np.nan_to_num(Y_coords_array)\n        \n        if np.isinf(X_geo_array).any():\n            print(f\"Found {np.isinf(X_geo_array).sum()} infinity values in X_geo. Replacing with large values.\")\n            X_geo_array = np.nan_to_num(X_geo_array, posinf=1e6, neginf=-1e6)\n            \n        if np.isinf(Y_coords_array).any():\n            print(f\"Found {np.isinf(Y_coords_array).sum()} infinity values in Y_coords. Replacing with large values.\")\n            Y_coords_array = np.nan_to_num(Y_coords_array, posinf=1e6, neginf=-1e6)\n        \n        print(f\"Generated {len(X_geo_array)} samples from {len(self.sequence)} sequences\")\n        truncated_count = sum([1 for m in metadata if m.get('truncated', False)])\n        print(f\"Truncated {truncated_count} sequences that were longer than {MAX_SEQ_LENGTH}\")\n        \n        return X_geo_array, Y_coords_array, self.normalization_params, metadata\n    \nclass MSAData(RNAData):\n    def __init__(self, input_sequences, input_labels=None, msa_folder=None, verbose=True, filtered_targets=None):\n        super(MSAData, self).__init__(input_sequences, input_labels)\n        self.msa_folder = msa_folder\n        self.verbose = verbose\n        self.msa_data = {}\n        \n        if filtered_targets is not None:\n            indices_to_keep = [i for i, tid in enumerate(self.target_id) if tid not in filtered_targets]\n            self.target_id = [self.target_id[i] for i in indices_to_keep]\n            self.sequence = [self.sequence[i] for i in indices_to_keep]\n            if self.verbose:\n                print(f\"MSAData: Filtered out {len(filtered_targets)} targets to maintain consistency\")\n    \n    def load_msa_files(self):\n        loaded = 0\n        skipped = 0\n        \n        if not self.msa_folder or not os.path.exists(self.msa_folder):\n            if self.verbose:\n                print(f\"MSA folder '{self.msa_folder}' doesn't exist\")\n            return self.msa_data\n            \n        for target_id in tqdm(self.target_id, desc=\"Loading MSA files\"):\n            if target_id in self.msa_data:  \n                pass\n                \n            possible_paths = [os.path.join(self.msa_folder, f\"{target_id}.MSA.fasta\")]\n            for msa_file_path in possible_paths:\n                if os.path.exists(msa_file_path) and os.path.getsize(msa_file_path) > 0:\n                    alignment = read_fasta(msa_file_path)\n                    if alignment and len(alignment) > 1:\n                        self.msa_data[target_id] = alignment\n                        loaded += 1\n                    else:\n                        skipped += 1\n        \n        if self.verbose:\n            print(f\"Loaded {loaded} MSA files, skipped {skipped}\")\n            if loaded == 0:\n                print(\"Warning: No MSA files were loaded. Check the MSA_FOLDER path and file format.\")\n                \n        return self.msa_data\n    \n    def prepare_msa_data(self, MAX_SEQ_LENGTH=None):\n        if MAX_SEQ_LENGTH is None:\n            MAX_SEQ_LENGTH = globals().get('MAX_SEQ_LENGTH', 500)\n            \n        self.load_msa_files() \n        all_msa_features = []\n        metadata = []\n    \n        feature_size = 2 + MAX_SEQ_LENGTH\n        \n        for target_id, seq in tqdm(zip(self.target_id, self.sequence), desc=\"Preparing MSA data\"):\n            truncated_seq, _ = self.truncate_sequence(seq, None, MAX_SEQ_LENGTH)\n            \n            padded_features = np.zeros((MAX_SEQ_LENGTH, feature_size))\n            \n            if target_id in self.msa_data:\n                try:\n                    alignment = self.msa_data[target_id]\n                    msa_features = self.process_alignment(\n                        alignment, seq, truncated_seq, MAX_SEQ_LENGTH\n                    )\n                    \n                    if msa_features.shape[1] != feature_size:\n                        if self.verbose:\n                            print(f\"Warning: Inconsistent feature size for {target_id}: {msa_features.shape[1]} vs {feature_size}\")\n                        if msa_features.shape[1] < feature_size:\n                            temp = np.zeros((msa_features.shape[0], feature_size))\n                            temp[:, :msa_features.shape[1]] = msa_features\n                            msa_features = temp\n                        else:\n                            msa_features = msa_features[:, :feature_size]\n                    \n                    valid_len = min(msa_features.shape[0], MAX_SEQ_LENGTH)\n                    padded_features[:valid_len, :] = msa_features[:valid_len, :]\n                except:\n                    pass\n            \n            all_msa_features.append(padded_features)\n            \n            metadata.append({\n                'target_id': target_id,\n                'seq_length': len(truncated_seq),\n                'original_seq_length': len(seq),\n                'truncated': len(seq) > MAX_SEQ_LENGTH\n            })\n            \n        if self.verbose:\n            print(f\"MSA features shape: {len(all_msa_features)} x {all_msa_features[0].shape[0]} x {all_msa_features[0].shape[1]}\")\n            print(f\"Generated {len(all_msa_features)} MSA samples from {len(self.sequence)} sequences\")\n            truncated_count = sum([1 for m in metadata if m.get('truncated', False)])\n            print(f\"Truncated {truncated_count} sequences that were longer than {MAX_SEQ_LENGTH}\")\n            \n        all_msa_features_array = np.array(all_msa_features)\n        if np.isnan(all_msa_features_array).any() or np.isinf(all_msa_features_array).any():\n            if self.verbose:\n                print(f\"Found {np.isnan(all_msa_features_array).sum()} NaN and {np.isinf(all_msa_features_array).sum()} infinity values in MSA features.\")\n            all_msa_features_array = np.nan_to_num(all_msa_features_array)\n                \n        return all_msa_features_array, metadata\n    \n    def process_alignment(self, alignment, full_seq, truncated_seq, MAX_SEQ_LENGTH):\n        alignment_length = alignment.get_alignment_length()\n        num_sequences = len(alignment)\n        seq_length = len(truncated_seq)\n        \n        features = np.zeros((seq_length, 2 + MAX_SEQ_LENGTH))\n    \n        target_idx = None\n        for i, record in enumerate(alignment):\n            if full_seq in str(record.seq).replace('-', ''):\n                target_idx = i\n                \n        if target_idx is None:\n            return features\n        \n        seq_to_align_map = {}\n        seq_pos = 0\n        target_seq = str(alignment[target_idx].seq)\n        \n        for align_pos, char in enumerate(target_seq):\n            if char != '-': \n                if seq_pos < len(full_seq):\n                    seq_to_align_map[seq_pos] = align_pos\n                    seq_pos += 1\n        \n        nucleotide_map = {'A': 0, 'C': 1, 'G': 2, 'T': 3, 'U': 3}\n        pssm = np.zeros((4, alignment_length))\n        \n        columns = []\n        for j in range(alignment_length):\n            column = [str(rec.seq)[j] for rec in alignment]\n            columns.append(column)\n            \n            for nucleotide in column:\n                if nucleotide.upper() in nucleotide_map:\n                    idx = nucleotide_map[nucleotide.upper()]\n                    pssm[idx, j] += 1\n        \n        column_sums = np.sum(pssm, axis=0)\n        column_sums[column_sums == 0] = 1 \n        frequencies = pssm / column_sums[np.newaxis, :]\n        \n        epsilon = 1e-10\n        entropy = -np.sum(frequencies * np.log2(frequencies + epsilon), axis=0)\n        max_entropy = -np.log2(0.25)\n        conservation_scores = 1 - (entropy / max_entropy)\n        conservation_scores = np.nan_to_num(conservation_scores)\n        \n        coverage = []\n        for j in range(alignment_length):\n            coverage.append(1.0 - columns[j].count('-') / num_sequences)\n        \n        truncated_align_positions = []\n        for pos in range(seq_length):\n            if pos in seq_to_align_map:\n                truncated_align_positions.append(seq_to_align_map[pos])\n        \n        correlation_matrix = {}\n        \n        for idx_i, align_pos_i in enumerate(truncated_align_positions):\n            if align_pos_i >= len(columns):\n                continue\n            col_i = columns[align_pos_i]\n            target_val_i = target_seq[align_pos_i]\n            \n            for idx_j, align_pos_j in enumerate(truncated_align_positions):\n                if idx_i != idx_j and align_pos_j < len(columns):\n                    col_j = columns[align_pos_j]\n                    target_val_j = target_seq[align_pos_j]\n                    \n                    matches = 0\n                    valid_pairs = 0\n                    \n                    for val_i, val_j in zip(col_i, col_j):\n                        if val_i != '-' and val_j != '-':\n                            valid_pairs += 1\n                            if (val_i == target_val_i) == (val_j == target_val_j):\n                                matches += 1\n                    \n                    corr = matches / valid_pairs if valid_pairs > 0 else 0\n                    correlation_matrix[(idx_i, idx_j)] = corr\n        \n        for seq_pos in range(seq_length):\n            if seq_pos in seq_to_align_map:\n                align_pos = seq_to_align_map[seq_pos]\n                \n                if align_pos < len(conservation_scores):\n                    features[seq_pos, 0] = conservation_scores[align_pos]\n                \n                features[seq_pos, 1] = coverage[align_pos] if align_pos < len(coverage) else 0\n                \n                if seq_pos < len(truncated_align_positions):\n                    for other_seq_pos in range(min(seq_length, MAX_SEQ_LENGTH)):\n                        if other_seq_pos in seq_to_align_map and seq_pos != other_seq_pos:\n                            if other_seq_pos < len(truncated_align_positions):\n                                features[seq_pos, 2 + other_seq_pos] = correlation_matrix.get((seq_pos, other_seq_pos), 0)\n        \n        return features\n\n\nLIMIT = 5\nAUGMENT = False\nNUM_AUG = 5\n\nif train_sequences is not None:\n    limited_train_sequences = train_sequences[:LIMIT]\n    limited_validation_sequences = validation_sequences[:LIMIT]\n    \n    geo_train = Geodata(limited_train_sequences, train_labels)\n    X_train_geo, Y_train_coords, normalization_params, train_metadata = geo_train.prepare_geo_data()\n    \n    geo_val = Geodata(limited_validation_sequences, validation_labels)\n    X_val_geo, Y_val_coords, _, _ = geo_val.prepare_geo_data()\n    \n    # filtered_train_targets = set([row[-1]['target_id'] for row in limited_train_sequences.iterrows()]) - set(geo_train.target_id)\n    # filtered_val_targets = set([row[-1]['target_id'] for row in limited_validation_sequences.iterrows()]) - set(geo_val.target_id)\n    \n    #X_train_msa, train_metadata = MSAData(limited_train_sequences, train_labels, MSA_FOLDER, filtered_targets=filtered_train_targets).prepare_msa_data()\n    #X_val_msa, _ = MSAData(limited_validation_sequences, validation_labels, MSA_FOLDER, filtered_targets=filtered_val_targets).prepare_msa_data()\n\nprint(\"X_train_geo: \", X_train_geo.shape)\n#print(\"X_train_msa: \", X_train_msa.shape)\nprint(\"Y_train: \", Y_train_coords.shape)\nprint(\"X_val_geo: \", X_val_geo.shape)\n#print(\"X_val_msa: \", X_val_msa.shape)\nprint(\"Y_val: \", Y_val_coords.shape)\nprint(\"variance of y train: \", np.var(Y_train_coords))\nprint(\"variance of y val: \", np.var(Y_val_coords))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T03:23:59.148915Z","iopub.execute_input":"2025-05-29T03:23:59.149821Z","iopub.status.idle":"2025-05-29T03:24:04.853775Z","shell.execute_reply.started":"2025-05-29T03:23:59.149738Z","shell.execute_reply":"2025-05-29T03:24:04.852379Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorflow as tf\nimport keras as keras\nfrom keras.src.saving import load_model\nfrom keras.src.models import Model, Sequential\nfrom keras.src.layers import Input, Layer, Dense, LayerNormalization, Dropout, MultiHeadAttention, GlobalAveragePooling1D, Lambda, Reshape\nfrom keras.src.optimizers import Adam\nfrom keras.src.metrics import Metric\nfrom keras.src.callbacks import EarlyStopping, ReduceLROnPlateau, TensorBoard, TerminateOnNaN","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T03:24:04.854953Z","iopub.execute_input":"2025-05-29T03:24:04.855265Z","iopub.status.idle":"2025-05-29T03:24:08.476384Z","shell.execute_reply.started":"2025-05-29T03:24:04.855224Z","shell.execute_reply":"2025-05-29T03:24:08.475077Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BATCH_SIZE = 1\nEPOCHS = 2\nLEARNING_RATE = 0.001\nTRANSFORMER_HEADS = 8\nDROPOUT_RATE = 0.2\nTRANSFORMER_UNITS = 256\nHIDDEN_DIM = 128\nFF_DIM = 512","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T03:24:08.477482Z","iopub.execute_input":"2025-05-29T03:24:08.478360Z","iopub.status.idle":"2025-05-29T03:24:08.484051Z","shell.execute_reply.started":"2025-05-29T03:24:08.478309Z","shell.execute_reply":"2025-05-29T03:24:08.482845Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@keras.src.saving.register_keras_serializable()\nclass TriangularAttention(Layer):\n    def __init__(self, hidden_dim=128, num_heads=4, dropout_rate=0.1, **kwargs):\n        super(TriangularAttention, self).__init__(**kwargs)\n        self.hidden_dim = hidden_dim\n        self.num_heads = num_heads\n        self.head_dim = hidden_dim // num_heads\n        self.dropout_rate = dropout_rate\n        \n    def build(self, input_shape):\n        self.attention = MultiHeadAttention(num_heads=self.num_heads, key_dim=self.head_dim, dropout=self.dropout_rate)\n        self.layernorm1 = LayerNormalization(epsilon=1e-6)\n        self.layernorm2 = LayerNormalization(epsilon=1e-6)\n        self.dropout = Dropout(self.dropout_rate)\n        \n        self.triangle_dense1 = Dense(self.hidden_dim, activation='relu')\n        self.triangle_dense2 = Dense(self.hidden_dim)\n        \n        super(TriangularAttention, self).build(input_shape)\n        \n    def call(self, inputs, training=False):\n        attn_output = self.attention(query=inputs, key=inputs, value=inputs, training=training)\n        attn_output = self.dropout(attn_output, training=training)\n        out1 = self.layernorm1(inputs + attn_output)\n        \n        seq_len = tf.shape(out1)[1]\n        \n        reshaped = tf.reshape(out1, [-1, seq_len, 1, self.hidden_dim])\n        repeated = tf.tile(reshaped, [1, 1, seq_len, 1])\n        \n        pairwise_features = tf.concat([repeated, tf.transpose(repeated, [0, 2, 1, 3])], axis=-1)\n        triangle_out = self.triangle_dense1(pairwise_features)\n        triangle_out = self.triangle_dense2(triangle_out)\n        triangle_out = tf.reduce_mean(triangle_out, axis=2)\n        output = self.layernorm2(out1 + triangle_out)\n        return output\n    \n\n\n@keras.saving.register_keras_serializable()\nclass CrossAttentionBlock(Layer):\n    def __init__(self, hidden_dim=128, num_heads=4, dropout_rate=0.1, **kwargs):\n        super(CrossAttentionBlock, self).__init__(**kwargs)\n        self.hidden_dim = hidden_dim\n        self.num_heads = num_heads\n        self.dropout_rate = dropout_rate\n        \n    def build(self, input_shape):\n        # bio to geo\n        self.cross_attention_1to2 = MultiHeadAttention(num_heads=self.num_heads, key_dim=self.hidden_dim // self.num_heads, dropout=self.dropout_rate)\n        \n        # geo to bio\n        self.cross_attention_2to1 = MultiHeadAttention(num_heads=self.num_heads, key_dim=self.hidden_dim // self.num_heads, dropout=self.dropout_rate)\n        \n        self.ffn_1 = Sequential([\n            Dense(self.hidden_dim * 4, activation='relu'), \n            Dropout(self.dropout_rate), \n            Dense(self.hidden_dim)\n        ])\n        \n        self.ffn_2 = Sequential([\n            Dense(self.hidden_dim * 4, activation='relu'), \n            Dropout(self.dropout_rate), \n            Dense(self.hidden_dim)\n        ])\n        \n        self.layernorm_1a = LayerNormalization(epsilon=1e-6)\n        self.layernorm_1b = LayerNormalization(epsilon=1e-6)\n        self.layernorm_2a = LayerNormalization(epsilon=1e-6)\n        self.layernorm_2b = LayerNormalization(epsilon=1e-6)\n        \n        self.dropout = Dropout(self.dropout_rate)\n        \n        super(CrossAttentionBlock, self).build(input_shape)\n        \n    def call(self, inputs, training=False):\n        tower1_input, tower2_input = inputs\n        \n        attn_1to2_output = self.cross_attention_1to2(query=tower1_input, key=tower2_input, value=tower2_input, training=training)\n        attn_1to2_output = self.dropout(attn_1to2_output, training=training)\n        tower1_output_temp = self.layernorm_1a(tower1_input + attn_1to2_output)\n        \n        ffn1_output = self.ffn_1(tower1_output_temp, training=training)\n        ffn1_output = self.dropout(ffn1_output, training=training)\n        tower1_output = self.layernorm_1b(tower1_output_temp + ffn1_output)\n        \n        attn_2to1_output = self.cross_attention_2to1(query=tower2_input, key=tower1_input, value=tower1_input, training=training)\n        attn_2to1_output = self.dropout(attn_2to1_output, training=training)\n        tower2_output_temp = self.layernorm_2a(tower2_input + attn_2to1_output)\n    \n        ffn2_output = self.ffn_2(tower2_output_temp, training=training)\n        ffn2_output = self.dropout(ffn2_output, training=training)\n        tower2_output = self.layernorm_2b(tower2_output_temp + ffn2_output)\n        \n        return tower1_output, tower2_output\n    \n\n@keras.saving.register_keras_serializable()\nclass BiologyTower(Layer):\n    def __init__(self, hidden_dim=128, num_layers=3, num_heads=4, dropout_rate=0.1, **kwargs):\n        super(BiologyTower, self).__init__(**kwargs)\n        self.hidden_dim = hidden_dim\n        self.num_layers = num_layers\n        self.num_heads = num_heads\n        self.dropout_rate = dropout_rate\n        \n    def build(self, input_shape):\n        self.input_projection = Dense(self.hidden_dim)\n        \n        self.transformer_blocks = [\n            TriangularAttention(\n                hidden_dim=self.hidden_dim, \n                num_heads=self.num_heads, \n                dropout_rate=self.dropout_rate, \n                name=f\"biology_triangular_attention_{i}\"\n            ) for i in range(self.num_layers)\n        ]\n        \n        self.layernorms = [LayerNormalization(epsilon=1e-6) for _ in range(self.num_layers)]\n        self.dropouts = [Dropout(self.dropout_rate) for _ in range(self.num_layers)]\n        \n        self.output_projection = Dense(self.hidden_dim)\n        \n        super(BiologyTower, self).build(input_shape)\n        \n    def call(self, inputs, training=False):\n        x = self.input_projection(inputs)\n    \n        for i in range(self.num_layers):\n            x = self.transformer_blocks[i](x, training=training)\n            x = self.dropouts[i](x, training=training)\n            x = self.layernorms[i](x)\n        \n        return self.output_projection(x)\n    \n\n@keras.saving.register_keras_serializable()\nclass GeometryTower(Layer):\n    def __init__(self, hidden_dim=128, num_layers=3, num_heads=4, dropout_rate=0.1, **kwargs):\n        super(GeometryTower, self).__init__(**kwargs)\n        self.hidden_dim = hidden_dim\n        self.num_layers = num_layers\n        self.num_heads = num_heads\n        self.dropout_rate = dropout_rate\n\n    def build(self, input_shape):\n        self.input_projection = Dense(self.hidden_dim)\n\n        self.transformer_blocks = [\n            TriangularAttention(\n                hidden_dim=self.hidden_dim, \n                num_heads=self.num_heads, \n                dropout_rate=self.dropout_rate, \n                name=f\"geometry_triangular_attention_{i}\"\n            ) for i in range(self.num_layers)\n        ]\n\n        self.layernorms = [LayerNormalization(epsilon=1e-6) for _ in range(self.num_layers)]\n        self.dropouts = [Dropout(self.dropout_rate) for _ in range(self.num_layers)]\n        \n        self.distance_mlp = Sequential([\n            Dense(self.hidden_dim * 2, activation='relu'), \n            Dropout(self.dropout_rate), \n            Dense(self.hidden_dim, activation='relu'), \n            Dense(1)\n        ])\n        \n        self.output_projection = Dense(self.hidden_dim)\n        \n        super(GeometryTower, self).build(input_shape)\n        \n    def call(self, inputs, training=False):\n        x = self.input_projection(inputs)\n        \n        for i in range(self.num_layers):\n            x = self.transformer_blocks[i](x, training=training)\n            x = self.dropouts[i](x, training=training)\n            x = self.layernorms[i](x)\n        \n        return self.output_projection(x)\n    \n    \n@keras.saving.register_keras_serializable()\nclass RNAStructurePredictor(Model):\n    def __init__(self, max_seq_length, msa_feature_dim, geo_feature_dim, \n                 hidden_dim=128, num_tower_layers=3, num_cross_layers=2, \n                 num_heads=4, dropout_rate=0.1, **kwargs):\n        super(RNAStructurePredictor, self).__init__(**kwargs)\n        \n        self.max_seq_length = max_seq_length\n        #self.msa_feature_dim = msa_feature_dim\n        self.geo_feature_dim = geo_feature_dim\n        self.hidden_dim = hidden_dim\n        self.num_tower_layers = num_tower_layers\n        self.num_cross_layers = num_cross_layers\n        self.num_heads = num_heads\n        self.dropout_rate = dropout_rate\n        \n    def build(self, input_shape=None):\n        # if input_shape is None:\n        #     msa_input_shape = (None, self.max_seq_length, self.msa_feature_dim)\n        #     geo_input_shape = (None, self.max_seq_length, self.geo_feature_dim)\n        # else:\n        #     msa_input_shape, geo_input_shape = input_shape\n        geo_input_shape = input_shape\n        \n        self.biology_tower = BiologyTower(\n            hidden_dim=self.hidden_dim, \n            num_layers=self.num_tower_layers, \n            num_heads=self.num_heads, \n            dropout_rate=self.dropout_rate\n        )\n        \n        self.geometry_tower = GeometryTower(\n            hidden_dim=self.hidden_dim, \n            num_layers=self.num_tower_layers, \n            num_heads=self.num_heads, \n            dropout_rate=self.dropout_rate\n        )\n        \n        self.cross_attention_blocks = [\n            CrossAttentionBlock(\n                hidden_dim=self.hidden_dim, \n                num_heads=self.num_heads, \n                dropout_rate=self.dropout_rate, \n                name=f\"cross_attention_block_{i}\"\n            ) for i in range(self.num_cross_layers)\n        ]\n        \n        self.coordinate_prediction = Sequential([\n            Dense(self.hidden_dim * 2, activation='relu'), \n            Dropout(self.dropout_rate), \n            Dense(self.hidden_dim, activation='relu'),  \n            Dropout(self.dropout_rate), \n            Dense(3)\n        ], name=\"coordinate_predictor\")\n        \n        super(RNAStructurePredictor, self).build(input_shape)\n    \n    def call(self, inputs, training=False):\n        #msa_input, geo_input = inputs\n        geo_input = inputs\n\n        #bio_features = self.biology_tower(msa_input, training=training)\n        geo_features = self.geometry_tower(geo_input, training=training)\n\n       #tower1_output, tower2_output = bio_features, geo_features\n        #for cross_block in self.cross_attention_blocks:\n         #   tower1_output, tower2_output = cross_block([tower1_output, tower2_output], training=training)\n        \n        coordinates = self.coordinate_prediction(geo_features, training=training)\n        return coordinates\n    \n\n#MSA_DIM = X_train_msa.shape[2]\nGEO_DIM = X_train_geo.shape[2]\n\nmodel = RNAStructurePredictor(MAX_SEQ_LENGTH, 14, GEO_DIM)\n\n#dummy_msa = tf.zeros((1, MAX_SEQ_LENGTH, X_train_msa.shape[2]))\ndummy_geo = tf.zeros((1, MAX_SEQ_LENGTH, X_train_geo.shape[2]))\noutput = model(dummy_geo)\nmodel.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T03:24:08.485522Z","iopub.execute_input":"2025-05-29T03:24:08.486370Z","iopub.status.idle":"2025-05-29T03:24:08.972419Z","shell.execute_reply.started":"2025-05-29T03:24:08.486313Z","shell.execute_reply":"2025-05-29T03:24:08.971308Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# class TMScore(Metric):\n#     def __init__(self, name='tm_score'):\n#         super(TMScore, self).__init__(name=name)\n#         self.tm_scores = self.add_weight(name='tm_scores', initializer='zeros')\n#         self.tm_count = self.add_weight(name='tm_count', initializer='zeros')\n\n#     def result(self):\n#         return tf.math.divide_no_nan(self.tm_scores, self.tm_count)\n    \n#     def reset_state(self):\n#         self.tm_scores.assign(0.0)\n#         self.tm_count.assign(0.0)\n\n#     def update_state(self, y_true, y_pred, sample_weight=None):\n#         batch_tm_scores = self.calculate_tm_scores(y_true, y_pred)\n#         batch_size = tf.cast(tf.shape(y_true)[0], dtype=tf.float32)\n#         self.tm_scores.assign_add(tf.reduce_sum(batch_tm_scores))\n#         self.tm_count.assign_add(batch_size)\n\n#     def calculate_tm_scores(self, y_true, y_pred):\n#         batch_size = tf.shape(y_true)[0]\n        \n#         def calculate_single_tm_score(structures):\n#             ref_structure, pred_structure = structures\n            \n#             L_ref = tf.cast(tf.shape(ref_structure)[0], tf.float32)\n#             d0 = self.calculate_d0(L_ref)\n#             squared_distances = tf.reduce_sum(tf.square(ref_structure - pred_structure), axis=1)\n#             tm_sum = tf.reduce_sum(1.0 / (1.0 + squared_distances / tf.square(d0)))\n#             tm_score = (1.0 / L_ref) * tm_sum\n#             return tm_score\n    \n#         return tf.map_fn(calculate_single_tm_score, (y_true, y_pred),fn_output_signature=tf.float32)\n    \n#     def calculate_d0(self, L_ref):\n#         d0_large = 0.6 * tf.sqrt(L_ref - 0.5) - 2.5\n        \n#         def d0_lt_12(): return tf.constant(0.3, dtype=tf.float32)\n#         def d0_12_15(): return tf.constant(0.4, dtype=tf.float32)\n#         def d0_16_19(): return tf.constant(0.5, dtype=tf.float32)\n#         def d0_20_23(): return tf.constant(0.6, dtype=tf.float32)\n#         def d0_24_29(): return tf.constant(0.7, dtype=tf.float32)\n        \n#         conditions = [(L_ref < 12, d0_lt_12), (tf.logical_and(L_ref >= 12, L_ref <= 15), d0_12_15), (tf.logical_and(L_ref >= 16, L_ref <= 19), d0_16_19), (tf.logical_and(L_ref >= 20, L_ref <= 23), d0_20_23), (tf.logical_and(L_ref >= 24, L_ref <= 29), d0_24_29)]\n#         d0_small = tf.case(conditions, default=lambda: tf.constant(0.0, dtype=tf.float32))\n#         d0 = tf.cond(L_ref >= 30, lambda: d0_large, lambda: d0_small)\n#         return d0\n\n        \n\n# @tf.function\n# def RMSD(y_true, y_pred):\n#     return tf.sqrt(tf.reduce_mean(tf.square(y_true - y_pred)))\n\noptimizer = Adam(learning_rate=LEARNING_RATE, clipnorm=1.0)\nloss_func = 'mse'\ntm_metric = 'mae' #TMScore(name='tm_score')\nmodel.compile(optimizer=optimizer, loss=loss_func, metrics=[tm_metric])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T03:24:08.973616Z","iopub.execute_input":"2025-05-29T03:24:08.974036Z","iopub.status.idle":"2025-05-29T03:24:08.991897Z","shell.execute_reply.started":"2025-05-29T03:24:08.973999Z","shell.execute_reply":"2025-05-29T03:24:08.990765Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"early_stop = EarlyStopping(monitor='val_loss', patience=5, restore_best_weights=True)\nreduce_lr = ReduceLROnPlateau(monitor='val_loss', factor=0.5, patience=3, min_lr=1e-6)\ntensorboard = TensorBoard(log_dir='logs')\nterminate_nan = TerminateOnNaN()\ncallbacks = [early_stop, reduce_lr, tensorboard, terminate_nan]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T03:24:08.994541Z","iopub.execute_input":"2025-05-29T03:24:08.994920Z","iopub.status.idle":"2025-05-29T03:24:09.006610Z","shell.execute_reply.started":"2025-05-29T03:24:08.994887Z","shell.execute_reply":"2025-05-29T03:24:09.005391Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if train_sequences is not None:\n    history = model.fit(X_train_geo, Y_train_coords, epochs=EPOCHS, validation_data=(X_val_geo, Y_val_coords), callbacks=callbacks, batch_size=BATCH_SIZE, validation_batch_size=BATCH_SIZE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T03:24:09.007946Z","iopub.execute_input":"2025-05-29T03:24:09.008283Z","iopub.status.idle":"2025-05-29T03:24:25.460786Z","shell.execute_reply.started":"2025-05-29T03:24:09.008243Z","shell.execute_reply":"2025-05-29T03:24:25.459449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def denormalize_coordinates(normalized_coords, normalization_params):\n    if normalized_coords.shape[0] == 0:\n        return normalized_coords\n    \n    if normalized_coords.ndim == 1:\n        normalized_coords = np.expand_dims(normalized_coords, axis=0)\n        \n    denormalized_coords = (normalized_coords * normalization_params['std']) + normalization_params['mean']\n    \n    return denormalized_coords\n\n\ndef create_dummy_submission_file(submission_path='submission.csv', num_predictions=5, num_sequences=10):\n    import random\n    import string\n    \n    columns = ['ID', 'resname', 'resid']\n    for k in range(1, num_predictions + 1):\n        columns.extend([f'x_{k}', f'y_{k}', f'z_{k}'])\n    \n    rows = []\n    nucleotides = ['A', 'U', 'G', 'C']\n    \n    for seq_idx in range(num_sequences):\n        target_id = ''.join(random.choices(string.ascii_uppercase + string.digits, k=8))\n        sequence_length = random.randint(10, 100)\n        sequence = ''.join(random.choices(nucleotides, k=sequence_length))\n        \n        for j, nucleotide in enumerate(sequence):\n            row_data = {'ID': f'{target_id}_{j+1}', 'resname': nucleotide, 'resid': j+1}\n        \n            for k in range(num_predictions):\n                row_data[f'x_{k+1}'] = float(np.random.uniform(-50, 50))\n                row_data[f'y_{k+1}'] = float(np.random.uniform(-50, 50))\n                row_data[f'z_{k+1}'] = float(np.random.uniform(-50, 50))\n            \n            rows.append(row_data)\n    \n    submission_df = pd.DataFrame(rows, columns=columns)\n    submission_df.to_csv(submission_path, index=False)\n    print(f\"Dummy submission file saved to {submission_path}\")\n    print(f\"Generated {len(submission_df)} rows from {num_sequences} dummy sequences\")\n    return submission_df\n\ndef create_submission_file(model, test_sequences, submission_path='submission.csv', num_predictions=5):\n    X_test_geo, _, _, _ = Geodata(test_sequences, None).prepare_geo_data()\n    X_test_msa, _ = MSAData(test_sequences, None, MSA_FOLDER).prepare_msa_data()\n    \n    columns = ['ID', 'resname', 'resid']\n    for k in range(1, num_predictions + 1):\n        columns.extend([f'x_{k}', f'y_{k}', f'z_{k}'])\n    \n    rows = []\n    \n    for i, (_, row) in enumerate(test_sequences.iterrows()):\n        sequence = row['sequence']\n        sequence_length = len(sequence)\n        target_id = row['target_id']\n        \n        if i >= len(X_test_geo) or i >= len(X_test_msa):\n            print(f\"Warning: Index {i} out of bounds for test data\")\n            continue\n            \n        sequence_geo_data = np.expand_dims(X_test_geo[i], axis=0)\n        sequence_msa_data = np.expand_dims(X_test_msa[i], axis=0)\n        \n        try:\n            #predictions = model.predict([sequence_msa_data, sequence_geo_data], verbose=0)\n            predictions = model.predict(sequence_geo_data, verbose=0)\n            \n            if predictions.ndim == 3:\n                predictions = np.squeeze(predictions, axis=0) \n            \n            predictions = np.nan_to_num(predictions, nan=0.0, posinf=1e6, neginf=-1e6)\n            \n            flat_coords = predictions[:sequence_length].flatten()\n            lower_bound = np.min(flat_coords)\n            higher_bound = np.max(flat_coords)\n            \n            for j, nucleotide in enumerate(sequence):\n                if j >= MAX_SEQ_LENGTH:\n                    coords = np.random.uniform(lower_bound, higher_bound, size=(3,))\n                else:   \n                    coords = predictions[j]\n                \n                row_data = {'ID': f'{target_id}_{j+1}', 'resname': nucleotide, 'resid': j+1}\n                \n                try:\n                    denormalized_coords = denormalize_coordinates(coords, normalization_params)\n                    denormalized_coords = np.squeeze(denormalized_coords)\n                    \n                    if denormalized_coords.shape == ():\n                        denormalized_coords = np.array([denormalized_coords, 0.0, 0.0])\n                    elif len(denormalized_coords) < 3:\n                        padding = np.zeros(3 - len(denormalized_coords))\n                        denormalized_coords = np.concatenate([denormalized_coords, padding])\n                    \n                    for k in range(num_predictions):\n                        noise = np.random.normal(scale=0.5, size=3)\n                        row_data[f'x_{k+1}'] = float(denormalized_coords[0] + noise[0])\n                        row_data[f'y_{k+1}'] = float(denormalized_coords[1] + noise[1])\n                        row_data[f'z_{k+1}'] = float(denormalized_coords[2] + noise[2])\n                        \n                except Exception as e:\n                    print(f\"Error processing coordinates for {target_id}_{j+1}: {e}\")\n                    for k in range(num_predictions):\n                        row_data[f'x_{k+1}'] = float(np.random.uniform(-10, 10))\n                        row_data[f'y_{k+1}'] = float(np.random.uniform(-10, 10))\n                        row_data[f'z_{k+1}'] = float(np.random.uniform(-10, 10))\n                \n                rows.append(row_data)\n                \n        except Exception as e:\n            print(f\"Error predicting for sequence {i} (target_id: {target_id}): {e}\")\n            for j, nucleotide in enumerate(sequence):\n                row_data = {'ID': f'{target_id}_{j+1}', 'resname': nucleotide, 'resid': j+1}\n                for k in range(num_predictions):\n                    row_data[f'x_{k+1}'] = float(np.random.uniform(-10, 10))\n                    row_data[f'y_{k+1}'] = float(np.random.uniform(-10, 10))\n                    row_data[f'z_{k+1}'] = float(np.random.uniform(-10, 10))\n                rows.append(row_data)\n    \n    submission_df = pd.DataFrame(rows, columns=columns)\n    submission_df.to_csv(submission_path, index=False)\n    print(f\"Submission file saved to {submission_path}\")\n    print(f\"Generated {len(submission_df)} rows\")\n    return submission_df\n\n\ntry:\n    if test_sequences is not None:\n        submission_df = create_submission_file(model, test_sequences, num_predictions=5)\n        print(\"Submission preview:\")\n        print(submission_df.head(10))\n        print(f\"Total rows: {len(submission_df)}\")\n        print(f\"submisison shape: {list(submission_df.shape)}\")\n    else:\n        create_dummy_submission_file()\nexcept Exception as e:\n    print(f\"Error creating submission file: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T03:24:25.462303Z","iopub.execute_input":"2025-05-29T03:24:25.462860Z","iopub.status.idle":"2025-05-29T03:24:28.374964Z","shell.execute_reply.started":"2025-05-29T03:24:25.462807Z","shell.execute_reply":"2025-05-29T03:24:28.373609Z"}},"outputs":[],"execution_count":null}]}