{"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":"gpu","dataSources":[{"sourceId":87793,"databundleVersionId":11553390,"sourceType":"competition"}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:54:48.891988Z","iopub.execute_input":"2025-03-25T15:54:48.892288Z","iopub.status.idle":"2025-03-25T15:54:50.929963Z","shell.execute_reply.started":"2025-03-25T15:54:48.892257Z","shell.execute_reply":"2025-03-25T15:54:50.911822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\n# Load the datasets\ntrain_sequences = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/train_sequences.csv')\ntrain_labels = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/train_labels.csv')\nvalidation_sequences = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/validation_sequences.csv')\nvalidation_labels = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/validation_labels.csv')\ntest_sequence = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:54:50.936904Z","iopub.execute_input":"2025-03-25T15:54:50.937268Z","iopub.status.idle":"2025-03-25T15:54:52.019132Z","shell.execute_reply.started":"2025-03-25T15:54:50.937247Z","shell.execute_reply":"2025-03-25T15:54:52.018223Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_sequences['sequence_length'] = train_sequences['sequence'].apply(len)\nvalidation_sequences['sequence_length'] = validation_sequences['sequence'].apply(len)\n\nplt.figure(figsize=(10, 6))\nsns.histplot(train_sequences['sequence_length'], bins=50, kde=True, label='Train Sequences', color='blue')\nsns.histplot(validation_sequences['sequence_length'], bins=50, kde=True, label='Validation Sequences', color='orange')\nplt.title('Distribution of Sequence Lengths')\nplt.xlabel('Sequence Length')\nplt.ylabel('Frequency')\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:54:52.020780Z","iopub.execute_input":"2025-03-25T15:54:52.021108Z","iopub.status.idle":"2025-03-25T15:54:52.492575Z","shell.execute_reply.started":"2025-03-25T15:54:52.021076Z","shell.execute_reply":"2025-03-25T15:54:52.491693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from collections import Counter\ndef nucleotide_composition(sequence):\n    return dict(Counter(sequence))\n\ntrain_sequences['nucleotide_composition'] = train_sequences['sequence'].apply(nucleotide_composition)\nvalidation_sequences['nucleotide_composition'] = validation_sequences['sequence'].apply(nucleotide_composition)\n\ntrain_nucleotide_counts = pd.DataFrame(train_sequences['nucleotide_composition'].tolist()).fillna(0).sum()\nvalidation_nucleotide_counts = pd.DataFrame(validation_sequences['nucleotide_composition'].tolist()).fillna(0).sum()\n\nplt.figure(figsize=(10, 6))\ntrain_nucleotide_counts.plot(kind='bar', color='blue', label='Train Sequences')\nvalidation_nucleotide_counts.plot(kind='bar', color='orange', label='Validation Sequences', alpha=0.7)\nplt.title('Nucleotide Composition')\nplt.xlabel('Nucleotide')\nplt.ylabel('Count')\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:54:52.493674Z","iopub.execute_input":"2025-03-25T15:54:52.494014Z","iopub.status.idle":"2025-03-25T15:54:52.768281Z","shell.execute_reply.started":"2025-03-25T15:54:52.493992Z","shell.execute_reply":"2025-03-25T15:54:52.767496Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_sequences['temporal_cutoff'] = pd.to_datetime(train_sequences['temporal_cutoff'])\nvalidation_sequences['temporal_cutoff'] = pd.to_datetime(validation_sequences['temporal_cutoff'])\n\ntrain_sequences['year'] = train_sequences['temporal_cutoff'].dt.year\nvalidation_sequences['year'] = validation_sequences['temporal_cutoff'].dt.year\n\nplt.figure(figsize=(10, 6))\nsns.histplot(train_sequences['temporal_cutoff'], bins=50, kde=True, label='Train Sequences')\nsns.histplot(validation_sequences['temporal_cutoff'], bins=50, kde=True, label='Validation Sequences', color='orange')\nplt.title('Temporal Distribution of Sequences')\nplt.xlabel('Temporal Cutoff')\nplt.ylabel('Frequency')\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:54:52.768987Z","iopub.execute_input":"2025-03-25T15:54:52.769203Z","iopub.status.idle":"2025-03-25T15:54:53.256144Z","shell.execute_reply.started":"2025-03-25T15:54:52.769184Z","shell.execute_reply":"2025-03-25T15:54:53.255301Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Analyze the distribution of 3D coordinates in train_labels\ncoordinate_columns = [col for col in train_labels.columns if col.startswith(('x_', 'y_', 'z_'))]\ntrain_labels['num_structures'] = train_labels[coordinate_columns].count(axis=1) // 3\n\n# Plot number of structures per target\nplt.figure(figsize=(10, 6))\nsns.histplot(train_labels['num_structures'], bins=20, kde=True)\nplt.title('Distribution of Number of Structures per Target')\nplt.xlabel('Number of Structures')\nplt.ylabel('Frequency')\nplt.show()\n\n# Plot distribution of coordinates\nplt.figure(figsize=(15, 5))\nfor i, coord in enumerate(['x_1', 'y_1', 'z_1']):\n    plt.subplot(1, 3, i+1)\n    sns.histplot(train_labels[coord].dropna(), bins=50, kde=True)\n    plt.title(f'Distribution of {coord}')\n    plt.xlabel(coord)\n    plt.ylabel('Frequency')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:54:53.256935Z","iopub.execute_input":"2025-03-25T15:54:53.257243Z","iopub.status.idle":"2025-03-25T15:54:56.286641Z","shell.execute_reply.started":"2025-03-25T15:54:53.257220Z","shell.execute_reply":"2025-03-25T15:54:56.285825Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"duplicate_sequences = train_sequences[train_sequences.duplicated('sequence', keep=False)]\nprint(f\"Number of duplicate sequences in train_sequences: {len(duplicate_sequences)}\")\n\n# Check for duplicate target_ids in train_labels\nduplicate_targets = train_labels[train_labels.duplicated('ID', keep=False)]\nprint(f\"Number of duplicate targets in train_labels: {len(duplicate_targets)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:54:56.287523Z","iopub.execute_input":"2025-03-25T15:54:56.287884Z","iopub.status.idle":"2025-03-25T15:54:56.312590Z","shell.execute_reply.started":"2025-03-25T15:54:56.287859Z","shell.execute_reply":"2025-03-25T15:54:56.311773Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Extract coordinates for a sample target\nsample_target = train_labels\nx = sample_target['x_1'].values\ny = sample_target['y_1'].values\nz = sample_target['z_1'].values\n\n# Plot 3D structure\nfig = plt.figure(figsize=(10, 8))\nax = fig.add_subplot(111, projection='3d')\nax.scatter(x, y, z, c='blue', marker='o')\nax.set_title('3D RNA Structure')\nax.set_xlabel('X')\nax.set_ylabel('Y')\nax.set_zlabel('Z')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:54:56.313580Z","iopub.execute_input":"2025-03-25T15:54:56.313907Z","iopub.status.idle":"2025-03-25T15:54:58.838585Z","shell.execute_reply.started":"2025-03-25T15:54:56.313875Z","shell.execute_reply":"2025-03-25T15:54:58.837808Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 检查描述中是否提及配体\ntrain_sequences['has_ligand'] = train_sequences['description'].str.contains('ligand', case=False)\n\n# 比较配体结合与未结合序列的序列长度\nplt.figure(figsize=(10, 6))\nsns.boxplot(x='has_ligand', y='sequence_length', data=train_sequences)\nplt.title('Sequence Length for Ligand-Bound vs. Unbound Sequences') #配体结合与未结合序列的长度\nplt.xlabel('Has Ligand') #是否有配体\nplt.ylabel('Sequence Length')   #序列长度\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:54:58.840909Z","iopub.execute_input":"2025-03-25T15:54:58.841154Z","iopub.status.idle":"2025-03-25T15:54:58.983864Z","shell.execute_reply.started":"2025-03-25T15:54:58.841132Z","shell.execute_reply":"2025-03-25T15:54:58.983030Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"训练数据全部有配体","metadata":{}},{"cell_type":"code","source":"from sklearn.feature_extraction.text import CountVectorizer\n# 将序列转换为k-mer计数\nvectorizer = CountVectorizer(analyzer='char', ngram_range=(3, 3))\nkmer_counts = vectorizer.fit_transform(train_sequences['sequence'])\nfrom sklearn.cluster import KMeans\n# 执行k-means聚类\nkmeans = KMeans(n_clusters=5, random_state=42)\nclusters = kmeans.fit_predict(kmer_counts)\n\n# 将聚类结果添加到数据框\ntrain_sequences['cluster'] = clusters\n\n# 绘制聚类分布图\nplt.figure(figsize=(10, 6))\nsns.countplot(x='cluster', data=train_sequences)\nplt.title('Sequence Clusters') # 序列聚类分布\nplt.xlabel('Cluster') #聚类\nplt.ylabel('Count') # 计数\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:54:58.984910Z","iopub.execute_input":"2025-03-25T15:54:58.985135Z","iopub.status.idle":"2025-03-25T15:54:59.741194Z","shell.execute_reply.started":"2025-03-25T15:54:58.985115Z","shell.execute_reply":"2025-03-25T15:54:59.740336Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.preprocessing import StandardScaler\nfrom sklearn.decomposition import PCA\ncoordinate_data = train_labels[['x_1', 'y_1', 'z_1']].dropna().values\n\n# 标准化坐标数据\nscaler = StandardScaler()\nscaled_coordinates = scaler.fit_transform(coordinate_data)\n\n# 执行PCA降维\npca = PCA(n_components=2)\npca_result = pca.fit_transform(scaled_coordinates)\n\n# 绘制前两个主成分\nplt.figure(figsize=(10, 6))\nplt.scatter(pca_result[:, 0], pca_result[:, 1], alpha=0.5)\nplt.scatter(pca_result[:, 0], pca_result[:, 1], alpha=0.5)\nplt.title('PCA on 3D Coordinates of RNA Sequences') # RNA序列的3D坐标PCA\nplt.xlabel('Principal Component 1')# 主成分1\nplt.ylabel('Principal Component 2')\nplt.colorbar(label='Target ID')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:54:59.742126Z","iopub.execute_input":"2025-03-25T15:54:59.742406Z","iopub.status.idle":"2025-03-25T15:55:00.672111Z","shell.execute_reply.started":"2025-03-25T15:54:59.742349Z","shell.execute_reply":"2025-03-25T15:55:00.671200Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":" # Check for missing values\nprint(\"Missing values in train_sequences:\")\nprint(train_sequences.isnull().sum())\n\nprint(\"\\nMissing values in train_labels:\")\nprint(train_labels.isnull().sum())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:55:00.672953Z","iopub.execute_input":"2025-03-25T15:55:00.673198Z","iopub.status.idle":"2025-03-25T15:55:00.693841Z","shell.execute_reply.started":"2025-03-25T15:55:00.673177Z","shell.execute_reply":"2025-03-25T15:55:00.693043Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load Data\ntrain_labels = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\")\ntrain_sequence = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_sequences.csv\")\nval_labels = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/validation_labels.csv\")\nval_sequence = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/validation_sequences.csv\")\ntest_sequence = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\")\n\n# Fill missing values\ntrain_labels.fillna(0, inplace=True)\nvalidation_labels.fillna(0, inplace=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:55:00.695027Z","iopub.execute_input":"2025-03-25T15:55:00.695339Z","iopub.status.idle":"2025-03-25T15:55:00.949152Z","shell.execute_reply.started":"2025-03-25T15:55:00.695312Z","shell.execute_reply":"2025-03-25T15:55:00.948460Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Sequence Encoding\nseq_dict = {'A': 1, 'C': 2, 'G': 3, 'U': 4}\ndef seq_map(seq):\n    return [seq_dict.get(char, 0) for char in seq]\n\ntrain_sequence['encoded_seq'] = train_sequence['sequence'].apply(seq_map)\ntest_sequence['encoded_seq'] = test_sequence['sequence'].apply(seq_map)\nval_sequence['encoded_seq'] = val_sequence['sequence'].apply(seq_map)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:55:00.949951Z","iopub.execute_input":"2025-03-25T15:55:00.950254Z","iopub.status.idle":"2025-03-25T15:55:00.968612Z","shell.execute_reply.started":"2025-03-25T15:55:00.950224Z","shell.execute_reply":"2025-03-25T15:55:00.968043Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.layers import (\n    Input, Embedding, Dense, Conv1D, BatchNormalization,\n    Concatenate, UpSampling1D, MaxPooling1D, Layer, Dropout  \n)\nfrom tensorflow.keras.preprocessing.sequence import pad_sequences\nfrom sklearn.metrics import mean_squared_error\ndef generate_label_coord(df):\n    result = {}\n    df[\"label\"] = df.ID.str.rsplit('_', n=1, expand=True).iloc[:,0]\n    for _, row in df.iterrows():\n        label = row['label']\n        resid = row['resid']\n        if label not in result:\n            result[label] = []\n        \n        # 检查是否所有必需的列都存在\n        if all(col in row for col in ['x_1', 'y_1', 'z_1']):\n            # 如果只有 x_1, y_1, z_1 可用，则将它们复制为 x_2, y_2, z_2 等\n            coords = np.array([\n                [row['x_1'], row['y_1'], row['z_1']],\n                [row['x_1'], row['y_1'], row['z_1']],  # 复制为 x_2, y_2, z_2\n                [row['x_1'], row['y_1'], row['z_1']],  # 复制为 x_3, y_3, z_3\n                [row['x_1'], row['y_1'], row['z_1']],  # 复制为 x_4, y_4, z_4\n                [row['x_1'], row['y_1'], row['z_1']]   # 复制为 x_5, y_5, z_5\n            ], dtype=np.float32)\n        else:\n            # 如果没有坐标可用，则使用零\n            # coords = np.zeros(5, 3, dtype=np.float32)\n            coords = np.zeros((5, 3), dtype=np.float32)\n        result[label].append((resid, coords))\n    \n    for key in result:\n        coords = np.stack([c for r, c in result[key]])\n        result[key] = coords\n    \n    return result\n\ntrain_stacked_coords = generate_label_coord(train_labels)\nval_stacked_coords = generate_label_coord(val_labels)\n\ndef generate_dataset(seq, stacked_coords):\n    X, y, tids = [], [], []\n    for idx, row in seq.iterrows():\n        tid = row['target_id']\n        if tid in stacked_coords:\n            X.append(row['encoded_seq'])\n            y.append(stacked_coords[tid])\n            tids.append(tid)\n    return X, y, tids\n\ntrain_X, train_y, train_tids = generate_dataset(train_sequence, train_stacked_coords)\nval_X, val_y, val_tids = generate_dataset(val_sequence, val_stacked_coords)\n\n# 填充序列和坐标\nmax_len = max(len(seq) for seq in train_X)\ntrain_X_pad = pad_sequences(train_X, maxlen=max_len, padding='post', value=0)\nval_X_pad = pad_sequences(val_X, maxlen=max_len, padding='post', value=0)\ntest_X = test_sequence['encoded_seq'].tolist()\ntest_X_pad = pad_sequences(test_X, maxlen=max_len, padding='post', value=0)\n\ndef pad_coords(coords, max_len):\n    L = coords.shape[0]\n    if L < max_len:\n        pad_width = ((0, max_len-L), (0, 0), (0, 0))\n        return np.pad(coords, pad_width, mode='constant', constant_values=0)\n    else:\n        return coords\n\ntrain_y_pad = np.array([pad_coords(y, max_len) for y in train_y])\nval_y_pad = np.array([pad_coords(y, max_len) for y in val_y])\n\n# CNN模型\ndef build_cnn_model(max_len):\n    input_seq = Input(shape=(max_len,), name='input_seq')\n    x = Embedding(input_dim=5, output_dim=16, mask_zero=False, name='embedding')(input_seq)\n    x = Conv1D(filters=64, kernel_size=3, padding='same', activation='relu', name='conv1')(x)\n    x = BatchNormalization(name='norm1')(x)\n    x = Dropout(0.2, name='drop1')(x)\n    x = Conv1D(filters=64, kernel_size=3, padding='same', activation='relu', name='conv2')(x)\n    x = BatchNormalization(name='norm2')(x)\n    x = Dropout(0.2, name='drop2')(x)\n    # 每个残基输出15个值（5组x、y、z坐标）\n    x = Conv1D(filters=15, kernel_size=1, padding='same', activation='linear', name='predicted_coords')(x)\n    model = Model(inputs=input_seq, outputs=x)\n    model.compile(optimizer='adam', loss='mae')\n    return model\n\nclass BasicUNet(Layer):\n    \"\"\"1D版简化UNet替代原DiffusionLayer\"\"\"\n    def __init__(self, filters=64, **kwargs):\n        super().__init__(**kwargs)\n        # 编码器\n        self.conv1 = Conv1D(filters, 3, padding='same', activation='relu')\n        self.pool1 = MaxPooling1D(2)\n        self.conv2 = Conv1D(filters*2, 3, padding='same', activation='relu')\n        \n        # 解码器\n        self.up1 = UpSampling1D(2)\n        self.deconv1 = Conv1D(filters, 3, padding='same', activation='relu')\n        \n        # 跳跃连接\n        self.concat = Concatenate(axis=-1)\n        \n        # 最终输出\n        self.final_conv = Conv1D(filters, 1, activation='relu')\n\n    def call(self, inputs):\n        # 编码路径\n        x1 = self.conv1(inputs)\n        p1 = self.pool1(x1)\n        x2 = self.conv2(p1)\n        \n        # 解码路径\n        u1 = self.up1(x2)\n        u1 = self.deconv1(u1)\n        \n        # 跳跃连接\n        c1 = self.concat([u1, x1])\n        \n        # 最终输出\n        return self.final_conv(c1)\n\ndef build_diffusion_model(max_len):\n    input_seq = Input(shape=(max_len,), name='input_seq')\n    \n    # 嵌入层保持不变\n    x = Embedding(input_dim=5, output_dim=16, mask_zero=False)(input_seq)\n    \n    # 用BasicUNet替换原DiffusionLayer\n    x = BasicUNet(filters=64)(x)  # 主要修改点\n    \n    # 后续层保持原结构\n    x = Dropout(0.2)(x)\n    x = Dense(64, activation='relu')(x)\n    x = Dropout(0.2)(x)\n    x = Dense(15, activation='linear')(x)\n    \n    model = Model(inputs=input_seq, outputs=x)\n    model.compile(optimizer='adam', loss='mae')\n    return model\n\n# 训练CNN\ncnn_model = build_cnn_model(max_len)\ncnn_history = cnn_model.fit(\n    train_X_pad, train_y_pad.reshape(train_y_pad.shape[0], train_y_pad.shape[1], -1),\n    validation_data=(val_X_pad, val_y_pad.reshape(val_y_pad.shape[0], val_y_pad.shape[1], -1)),\n    epochs=50, batch_size=16, verbose=1\n)\n\n\n# 训练扩散模型\ndiffusion_model = build_diffusion_model(max_len)\ndiffusion_history = diffusion_model.fit(\n    train_X_pad, train_y_pad.reshape(train_y_pad.shape[0], train_y_pad.shape[1], -1),\n    validation_data=(val_X_pad, val_y_pad.reshape(val_y_pad.shape[0], val_y_pad.shape[1], -1)),\n    epochs=50, batch_size=16, verbose=1\n)\n\ndef evaluate_model(model, X, y):\n    preds = model.predict(X)\n    preds = preds.reshape(preds.shape[0], preds.shape[1], 5, 3)  # 重新形状为 [batch_size, seq_len, 5, 3]\n    rmse = np.sqrt(mean_squared_error(y.reshape(-1), preds.reshape(-1)))\n    return rmse\n\ncnn_rmse = evaluate_model(cnn_model, val_X_pad, val_y_pad)\ndiffusion_rmse = evaluate_model(diffusion_model, val_X_pad, val_y_pad)\nprint(f\"CNN验证集RMSE: {cnn_rmse}\")\nprint(f\"扩散模型验证集RMSE: {diffusion_rmse}\")\n\ncnn_preds = cnn_model.predict(test_X_pad)\ndiffusion_preds = diffusion_model.predict(test_X_pad)\n\n# 集成预测（加权平均）\nensemble_preds = 0.7 * cnn_preds + 0.3 * diffusion_preds\nensemble_preds = ensemble_preds.reshape(ensemble_preds.shape[0], ensemble_preds.shape[1], 5, 3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T16:07:52.961390Z","iopub.execute_input":"2025-03-25T16:07:52.961725Z","iopub.status.idle":"2025-03-25T16:09:13.490530Z","shell.execute_reply.started":"2025-03-25T16:07:52.961686Z","shell.execute_reply":"2025-03-25T16:09:13.489814Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission_rows = []\nfor idx, row in test_sequence.iterrows():\n    target_id = row['target_id']\n    coords = ensemble_preds[idx]  # Shape: [sequence_length, 5, 3]\n    seq_length = len(row['encoded_seq'])\n    coords = coords[:seq_length, :, :]  # Trim to the actual sequence length\n    for i in range(seq_length):\n        x_coords = coords[i, :, 0]  # x_1, x_2, x_3, x_4, x_5\n        y_coords = coords[i, :, 1]  # y_1, y_2, y_3, y_4, y_5\n        z_coords = coords[i, :, 2]  # z_1, z_2, z_3, z_4, z_5\n        submission_rows.append({\n            'ID': f\"{target_id}_{i+1}\",\n            'resname': row['sequence'][i],\n            'resid': i+1,\n            'x_1': x_coords[0], 'x_2': x_coords[1], 'x_3': x_coords[2], 'x_4': x_coords[3], 'x_5': x_coords[4],\n            'y_1': y_coords[0], 'y_2': y_coords[1], 'y_3': y_coords[2], 'y_4': y_coords[3], 'y_5': y_coords[4],\n            'z_1': z_coords[0], 'z_2': z_coords[1], 'z_3': z_coords[2], 'z_4': z_coords[3], 'z_5': z_coords[4]})\n\nsubmission = pd.DataFrame(submission_rows)\nsubmission.to_csv(\"submission.csv\", index=False)\nprint(\"Submission file created: submission.csv\") ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T16:09:56.758956Z","iopub.execute_input":"2025-03-25T16:09:56.759248Z","iopub.status.idle":"2025-03-25T16:09:56.830502Z","shell.execute_reply.started":"2025-03-25T16:09:56.759225Z","shell.execute_reply":"2025-03-25T16:09:56.829842Z"}},"outputs":[],"execution_count":null}]}