{"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":11403143,"sourceType":"competition"}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Stanford RNA 3D Folding Competition\n# Revised Notebook for RNA 3D Structure Prediction (Improved V3)\n\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport os\nfrom scipy.spatial.distance import cdist\nimport warnings\nimport random\nimport torch.nn.functional as F\nwarnings.filterwarnings('ignore')\n\nprint(\"Starting Stanford RNA 3D Folding notebook...\")\n\n# 1. Data Loading and Exploration\nprint(\"Loading 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_sequences = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/test_sequences.csv')\nsample_submission = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/sample_submission.csv')\n\n# Fill missing coordinate values to avoid NaNs.\ntrain_labels.fillna(0, inplace=True)\nvalidation_labels.fillna(0, inplace=True)\n\nprint(\"\\nBasic dataset information:\")\nprint(f\"Training sequences: {train_sequences.shape}\")\nprint(f\"Training labels: {train_labels.shape}\")\nprint(f\"Validation sequences: {validation_sequences.shape}\")\nprint(f\"Validation labels: {validation_labels.shape}\")\nprint(f\"Test sequences: {test_sequences.shape}\")\nprint(f\"Sample submission: {sample_submission.shape}\")\n\n# 2. Data Analysis (optional visualizations)\nprint(\"\\nAnalyzing RNA sequence lengths...\")\ntrain_sequences['length'] = train_sequences['sequence'].str.len()\nplt.figure(figsize=(12, 6))\nsns.histplot(train_sequences['length'], bins=50)\nplt.title('Distribution of RNA Sequence Lengths')\nplt.xlabel('Sequence Length')\nplt.ylabel('Count')\nplt.savefig('sequence_length_distribution.png')\nplt.close()\n\n# 3. Data Preprocessing\ndef preprocess_sequence_data(sequences_df, labels_df=None, is_train=True):\n    \"\"\"\n    Preprocess RNA sequence data.\n    Convert sequences to numerical form and normalize coordinate targets per sequence.\n    \"\"\"\n    nucleotide_map = {'A': 0, 'C': 1, 'G': 2, 'U': 3, 'T': 3}\n    processed_data = []\n    \n    for idx, row in sequences_df.iterrows():\n        seq_id = row['target_id']\n        sequence = row['sequence']\n        numerical_seq = [nucleotide_map.get(nuc, 4) for nuc in sequence]\n        \n        structures = None\n        if is_train and labels_df is not None:\n            sequence_labels = labels_df[labels_df['ID'].str.startswith(seq_id + '_')]\n            if not sequence_labels.empty:\n                num_structures = (len(sequence_labels.columns) - 3) // 3\n                structures = []\n                for i in range(1, num_structures + 1):\n                    coords = []\n                    for _, label_row in sequence_labels.iterrows():\n                        x = label_row[f'x_{i}']\n                        y = label_row[f'y_{i}']\n                        z = label_row[f'z_{i}']\n                        coords.append([x, y, z])\n                    coords = np.array(coords)\n                    # Normalize coordinates per sequence (center and scale)\n                    mean = np.mean(coords, axis=0)\n                    std = np.std(coords, axis=0) + 1e-8\n                    coords_norm = (coords - mean) / std\n                    structures.append(coords_norm)\n        processed_data.append({\n            'id': seq_id,\n            'sequence': numerical_seq,\n            'structures': structures\n        })\n    return processed_data\n\nprint(\"Preprocessing training data...\")\ntrain_data = preprocess_sequence_data(train_sequences, train_labels)\nprint(\"Preprocessing validation data...\")\nvalidation_data = preprocess_sequence_data(validation_sequences, validation_labels)\nprint(\"Preprocessing test data...\")\ntest_data = preprocess_sequence_data(test_sequences, is_train=False)\n\n# 4. Feature Engineering\ndef extract_sequence_features(sequence):\n    \"\"\"\n    Extract one-hot encoding, positional encoding, and GC-content as features.\n    \"\"\"\n    one_hot = np.zeros((len(sequence), 5))\n    for i, nucleotide in enumerate(sequence):\n        one_hot[i, nucleotide] = 1\n    gc_content = []\n    window_size = 5\n    for i in range(len(sequence)):\n        start = max(0, i - window_size // 2)\n        end = min(len(sequence), i + window_size // 2 + 1)\n        window = sequence[start:end]\n        gc_count = sum(1 for n in window if n in [1, 2])\n        gc_content.append(gc_count / len(window))\n    positions = np.array([[i / len(sequence)] for i in range(len(sequence))])\n    features = np.hstack((one_hot, positions, np.array(gc_content).reshape(-1, 1)))\n    return features\n\nprint(\"Extracting sequence features...\")\nfor i, data in enumerate(train_data):\n    train_data[i]['features'] = extract_sequence_features(data['sequence'])\nfor i, data in enumerate(validation_data):\n    validation_data[i]['features'] = extract_sequence_features(data['sequence'])\nfor i, data in enumerate(test_data):\n    test_data[i]['features'] = extract_sequence_features(data['sequence'])\n\n# 5. RNA Secondary Structure Prediction (simple rule-based)\ndef predict_rna_secondary_structure(sequence):\n    nucleotide_map_inv = {0: 'A', 1: 'C', 2: 'G', 3: 'U', 4: 'X'}\n    seq_chars = [nucleotide_map_inv[n] for n in sequence]\n    structure = ['.' for _ in range(len(seq_chars))]\n    complementary = {'A': 'U', 'U': 'A', 'G': 'C', 'C': 'G', 'X': None}\n    for i in range(len(seq_chars)):\n        if structure[i] != '.':\n            continue\n        for j in range(len(seq_chars) - 1, i + 3, -1):\n            if structure[j] != '.':\n                continue\n            if complementary[seq_chars[i]] == seq_chars[j]:\n                structure[i] = '('\n                structure[j] = ')'\n                break\n    return ''.join(structure)\n\ndef enhance_features_with_ss(data):\n    for i, item in enumerate(data):\n        seq = item['sequence']\n        ss = predict_rna_secondary_structure(seq)\n        ss_features = np.zeros((len(ss), 3))\n        for j, char in enumerate(ss):\n            if char == '.':\n                ss_features[j, 0] = 1\n            elif char == '(':\n                ss_features[j, 1] = 1\n            elif char == ')':\n                ss_features[j, 2] = 1\n        data[i]['features'] = np.hstack((item['features'], ss_features))\n    return data\n\nprint(\"Enhancing features with secondary structure information...\")\ntrain_data = enhance_features_with_ss(train_data)\nvalidation_data = enhance_features_with_ss(validation_data)\ntest_data = enhance_features_with_ss(test_data)\n\n\n\nclass RNADataset(Dataset):\n    def __init__(self, data, augment=False):\n        self.data = data\n        self.augment = augment\n        \n    def __len__(self): \n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        item = self.data[idx]\n        features = item['features']\n        \n        # Data augmentation\n        if self.augment and random.random() < 0.7:\n            if random.random() < 0.1:\n                mut_idx = random.randint(0, len(features)-1)\n                features[mut_idx,:5] = np.eye(5)[random.choice([0,1,2,3])]\n            if len(features) > 20 and random.random() < 0.3:\n                start = random.randint(0, len(features)-10)\n                features[start:start+5] = features[start:start+5][::-1]\n        \n        features = torch.tensor(features, dtype=torch.float32)\n        target = torch.tensor(item['structures'][0], dtype=torch.float32) if item['structures'] else None\n        return {\n            'features': features,\n            'target': target,\n            'length': features.shape[0],  # 直接使用整数长度\n            'id': item['id']             # 返回序列ID\n        }\n\ndef collate_fn(batch):\n    # 按序列长度排序\n    sorted_batch = sorted(batch, key=lambda x: x['length'], reverse=True)\n    \n    # 提取各组件\n    features = [x['features'] for x in sorted_batch]\n    targets = [x['target'] for x in sorted_batch]\n    lengths = [x['length'] for x in sorted_batch]\n    ids = [x['id'] for x in sorted_batch]\n    \n    # 填充特征\n    max_len = features[0].shape[0]\n    feat_dim = features[0].shape[1]\n    padded_features = torch.zeros((len(features), max_len, feat_dim))\n    for i, feat in enumerate(features):\n        padded_features[i, :len(feat)] = feat\n    \n    # 填充目标（如果有）\n    if targets[0] is not None:\n        padded_targets = torch.zeros((len(targets), max_len, 3))\n        for i, tgt in enumerate(targets):\n            padded_targets[i, :len(tgt)] = tgt\n    else:\n        padded_targets = None\n    \n    return {\n        'features': padded_features,\n        'targets': padded_targets,\n        'lengths': torch.tensor(lengths),  # 转换为tensor\n        'ids': ids\n    }\n\ntrain_dataset = RNADataset(train_data, augment=True)\nvalidation_dataset = RNADataset(validation_data)\ntest_dataset = RNADataset(test_data)\n\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=4,\n    shuffle=True,\n    collate_fn=collate_fn,\n    pin_memory=True,\n    num_workers=2,\n    persistent_workers=True  # 防止多epoch数据重载\n)\n\nvalidation_loader = DataLoader(\n    validation_dataset,\n    batch_size=4,\n    shuffle=True,\n    collate_fn=collate_fn,\n    pin_memory=True,\n    num_workers=2,\n    persistent_workers=True  # 防止多epoch数据重载\n)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=4,  # 减少测试批次大小\n    collate_fn=collate_fn,\n    pin_memory=True,\n    num_workers=2,\n    persistent_workers=True\n)\n\n# 3. Enhanced Model Architecture\nclass RNAFoldingModel(nn.Module):\n    def __init__(self, input_dim, hidden_dim=512, num_layers=4):\n        super().__init__()\n        \n        # BiLSTM Encoder\n        self.lstm = nn.LSTM(\n            input_dim, hidden_dim,\n            num_layers=num_layers,\n            bidirectional=True,\n            batch_first=True,\n            dropout=0.3 if num_layers>1 else 0\n        )\n        \n        # Transformer Encoder\n        self.transformer = nn.TransformerEncoder(\n            nn.TransformerEncoderLayer(\n                d_model=2*hidden_dim,\n                nhead=8,\n                dim_feedforward=1024,\n                dropout=0.3\n            ),\n            num_layers=2\n        )\n        \n        # Geometric Attention\n        self.attention = nn.MultiheadAttention(\n            2*hidden_dim, num_heads=8, dropout=0.3\n        )\n        \n        # Dynamic Convolution\n        self.conv = nn.Sequential(\n            nn.Conv1d(2*hidden_dim, 512, kernel_size=5, padding=2),\n            nn.BatchNorm1d(512),\n            nn.GELU(),\n            nn.Conv1d(512, 256, kernel_size=3, padding=1),\n            nn.BatchNorm1d(256),\n            nn.GELU()\n        )\n        \n        # Prediction Head\n        self.head = nn.Sequential(\n            nn.Linear(256, 128),\n            nn.LayerNorm(128),\n            nn.GELU(),\n            nn.Linear(128, 3)\n        )\n        \n    def forward(self, x, lengths):\n        # BiLSTM\n        if isinstance(lengths, torch.Tensor):\n            lengths = lengths.cpu().numpy().tolist()\n        elif isinstance(lengths, list):\n            pass  # 已经是列表形式\n        else:\n            raise ValueError(f\"Unsupported lengths type: {type(lengths)}\")\n        \n        # BiLSTM处理\n        packed = nn.utils.rnn.pack_padded_sequence(\n            x,\n            lengths=lengths,\n            batch_first=True,\n            enforce_sorted=False\n        )\n        lstm_out, _ = self.lstm(packed)\n        lstm_out, _ = nn.utils.rnn.pad_packed_sequence(lstm_out, batch_first=True)\n        \n        # Transformer\n        transformer_out = self.transformer(lstm_out)\n        \n        # Attention\n        attn_out, _ = self.attention(\n            transformer_out.permute(1,0,2),\n            transformer_out.permute(1,0,2),\n            transformer_out.permute(1,0,2)\n        )\n        attn_out = attn_out.permute(1,0,2)\n        \n        # Residual Connection\n        combined = transformer_out + attn_out\n        \n        # Convolution\n        conv_out = self.conv(combined.permute(0,2,1)).permute(0,2,1)\n        \n        # Prediction\n        return self.head(conv_out)\n\n# 4. Enhanced Loss Function\nclass GeometricLoss(nn.Module):\n    def __init__(self, alpha=0.7):\n        super().__init__()\n        self.alpha = alpha\n        self.coord_loss = nn.SmoothL1Loss()\n        \n    def distance_loss(self, pred, target):\n        pred_dist = torch.cdist(pred, pred)\n        target_dist = torch.cdist(target, target)\n        return F.mse_loss(pred_dist, target_dist)\n    \n    def forward(self, pred, target, lengths):\n        total_loss = 0\n        for i, l in enumerate(lengths):\n            if l < 2: continue\n            pred_i = pred[i,:l]\n            target_i = target[i,:l]\n            \n            coord_loss = self.coord_loss(pred_i, target_i)\n            dist_loss = self.distance_loss(pred_i, target_i)\n            total_loss += self.alpha*coord_loss + (1-self.alpha)*dist_loss\n        return total_loss / len(lengths)\n\n# 5. Enhanced Training Loop\ndef train_model(model, train_loader, val_loader, epochs=50, lr=1e-4, device='cpu'):\n    device = torch.device(device)\n    model = model.to(device)\n    \n    optimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)\n    criterion = GeometricLoss(alpha=0.5)\n    \n    best_tm = 0.0\n    history = {'train_loss': [], 'val_loss': [], 'tm_score': []}\n    \n    for epoch in range(epochs):\n        # Training phase\n        model.train()\n        train_loss = 0.0\n        for batch in train_loader:\n            features = batch['features'].to(device)\n            targets = batch['targets'].to(device)\n            lengths = batch['lengths']\n            \n            optimizer.zero_grad()\n            outputs = model(features, lengths)\n            loss = criterion(outputs, targets, lengths)\n            loss.backward()\n            nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            \n            train_loss += loss.item()\n        \n        # Validation phase\n        model.eval()\n        val_loss = 0.0\n        tm_scores = []\n        with torch.no_grad():\n            for batch in val_loader:\n                features = batch['features'].to(device)\n                targets = batch['targets'].to(device)\n                lengths = batch['lengths']\n                ids = batch['ids']\n                \n                # 计算验证损失\n                outputs = model(features, lengths)\n                loss = criterion(outputs, targets, lengths)\n                val_loss += loss.item()\n                \n                # 计算TM-Score\n                outputs_np = outputs.cpu().numpy()\n                targets_np = targets.cpu().numpy()\n                for i in range(len(lengths)):\n                    l = lengths[i].item()\n                    if l < 5:  # 跳过过短序列\n                        continue\n                    pred_coords = outputs_np[i, :l, :]\n                    true_coords = targets_np[i, :l, :]\n                    tm = calculate_tm_score(pred_coords, true_coords)\n                    tm_scores.append(tm)\n        \n        # 记录指标\n        avg_train_loss = train_loss / len(train_loader)\n        avg_val_loss = val_loss / len(val_loader)\n        avg_tm = np.mean(tm_scores) if tm_scores else 0.0\n        \n        history['train_loss'].append(avg_train_loss)\n        history['val_loss'].append(avg_val_loss)\n        history['tm_score'].append(avg_tm)\n        \n        # 打印日志\n        print(f\"\\nEpoch {epoch+1}/{epochs}\")\n        print(f\"Train Loss: {avg_train_loss:.4f}\")\n        print(f\"Val Loss: {avg_val_loss:.4f}\")\n        print(f\"TM-Score: {avg_tm:.4f}\")\n        \n        # 保存最佳模型\n        if avg_tm > best_tm:\n            best_tm = avg_tm\n            torch.save(model.state_dict(), 'best_model.pth')\n    \n    return model\n\ndef train_epoch(model, dataloader, optimizer, device):\n    model.train()\n    epoch_loss = 0\n    batches = 0\n    # 修复点1: 正确的解包方式\n    for features, targets, seq_lengths in dataloader:\n        if targets is None:\n            continue\n        optimizer.zero_grad()\n        features = features.to(device)\n        targets = targets.to(device)\n        outputs = model(features, seq_lengths)\n        loss = GeometricLoss()(outputs, targets, seq_lengths)\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0)\n        optimizer.step()\n        epoch_loss += loss.item()\n        batches += 1\n    return epoch_loss / batches if batches > 0 else float('inf')\n\ndef validate(model, dataloader, device):\n    model.eval()\n    val_loss = 0\n    batches = 0\n    with torch.no_grad():\n        # 修复点2: 正确的解包方式\n        for features, targets, seq_lengths in dataloader:\n            if targets is None:\n                continue\n            features = features.to(device)\n            targets = targets.to(device)\n            outputs = model(features, seq_lengths)\n            loss = GeometricLoss()(outputs, targets, seq_lengths)\n            val_loss += loss.item()\n            batches += 1\n    return val_loss / batches if batches > 0 else float('inf')\n\ndef evaluate_model(model, dataloader, device):\n    model.eval()\n    tm_scores = []\n    with torch.no_grad():\n        # 修复点3: 正确的解包方式\n        for features, targets, seq_lengths in dataloader:\n            if targets is None:\n                continue\n            features = features.to(device)\n            outputs = model(features, seq_lengths)\n            outputs = outputs.cpu().numpy()\n            targets = targets.cpu().numpy()\n            for i, length in enumerate(seq_lengths):\n                pred_coords = outputs[i, :length, :]\n                target_coords = targets[i, :length, :]\n                tm_score = calculate_tm_score(pred_coords, target_coords)\n                tm_scores.append(tm_score)\n    return np.mean(tm_scores) if tm_scores else 0\n\n\ndef calculate_tm_score(predicted, reference):\n    l_ref = len(reference)\n    if l_ref < 5:  # 过短序列不计算\n        return 0.0\n    \n    # 标准d0计算公式\n    d0 = 1.24 * (l_ref - 15) ** (1/3) - 1.8\n    d0 = max(d0, 0.5)  # 确保最小值\n    \n    # 结构对齐（假设已预处理对齐）\n    aligned_len = min(len(predicted), l_ref)\n    pred = predicted[:aligned_len]\n    ref = reference[:aligned_len]\n    \n    # 计算距离矩阵\n    pred_dists = cdist(pred, pred)\n    ref_dists = cdist(ref, ref)\n    \n    # 计算TM-Score\n    tm = (1/(1 + ((pred_dists - ref_dists)/d0)**2)).sum()\n    tm_normalized = tm / (l_ref**2 - l_ref)  # 标准化\n    \n    return np.clip(tm_normalized, 0.0, 1.0)  # 强制限制在[0,1]\n\n\n\n# 9. Model Inference and Multiple Structure Generation\ndef generate_diverse_structures(model, features, seq_length, num_structures=5, noise_scale=0.05, device='cpu'):\n    model.eval()\n    structures = []\n    with torch.no_grad():\n        # 确保长度参数格式正确\n        if isinstance(seq_length, torch.Tensor):\n            seq_length = seq_length.item()\n            \n        for i in range(num_structures):\n            if i > 0:\n                noise = torch.randn_like(features) * noise_scale\n                perturbed_features = features + noise\n            else:\n                perturbed_features = features\n                \n            # 显式转换长度参数为张量\n            length_tensor = torch.tensor([seq_length], dtype=torch.long)\n            output = model(\n                perturbed_features.unsqueeze(0),\n                lengths=length_tensor.to(device)  # 保持设备一致\n            )\n            coords = output[0, :seq_length, :].cpu().numpy()\n            structures.append(coords)\n    return structures\n\ndef generate_predictions(model, dataloader, device, num_predictions=5):\n    model.eval()\n    all_predictions = {}\n    for batch in dataloader:\n        features = batch['features'].to(device)\n        lengths = batch['lengths'].tolist()  # 转换为列表\n        ids = batch['ids']\n        \n        for i in range(features.size(0)):\n            seq_id = ids[i]\n            seq_len = lengths[i]\n            \n            # 生成预测\n            seq_features = features[i, :seq_len, :]\n            predictions = generate_diverse_structures(\n                model,\n                seq_features,\n                seq_len,  # 直接使用整数值\n                num_structures=num_predictions,\n                device=device\n            )\n            all_predictions[seq_id] = predictions\n    return all_predictions\n\n# 10. Submission File Generation\n\ndef create_submission_file(predictions, test_sequences_df, output_file='submission.csv'):\n    submission_rows = []\n    for _, row in test_sequences_df.iterrows():\n        seq_id = row['target_id']\n        sequence = row['sequence']\n        if seq_id in predictions:\n            pred_structures = predictions[seq_id]\n            num_structures = len(pred_structures)\n            for i in range(len(sequence)):\n                submission_row = {\n                    'ID': f\"{seq_id}_{i+1}\",\n                    'resname': sequence[i],\n                    'resid': i+1\n                }\n                for j in range(5):\n                    if j < num_structures:\n                        coords = pred_structures[j][i]\n                        submission_row[f'x_{j+1}'] = coords[0]\n                        submission_row[f'y_{j+1}'] = coords[1]\n                        submission_row[f'z_{j+1}'] = coords[2]\n                    else:\n                        submission_row[f'x_{j+1}'] = submission_row[f'x_{j}']\n                        submission_row[f'y_{j+1}'] = submission_row[f'y_{j}']\n                        submission_row[f'z_{j+1}'] = submission_row[f'z_{j}']\n                submission_rows.append(submission_row)\n    submission_df = pd.DataFrame(submission_rows)\n    \n    # 确保输出目录存在\n    output_dir = os.path.dirname(output_file)\n    if output_dir and not os.path.exists(output_dir):\n        os.makedirs(output_dir, exist_ok=True)\n    \n    try:\n        submission_df.to_csv(output_file, index=False)\n        print(f\"文件已成功保存至: {os.path.abspath(output_file)}\")\n    except Exception as e:\n        print(f\"保存文件时出错: {e}\")\n        return None\n    \n    return submission_df\n\n# 11. Visualization Functions\ndef visualize_3d_structure(coords, title=\"RNA 3D Structure\"):\n    import matplotlib.pyplot as plt\n    fig = plt.figure(figsize=(10, 8))\n    ax = fig.add_subplot(111, projection='3d')\n    ax.scatter(coords[:, 0], coords[:, 1], coords[:, 2], c='blue', marker='o', s=30, label=\"C1' atoms\")\n    for i in range(len(coords) - 1):\n        ax.plot([coords[i, 0], coords[i+1, 0]], \n                [coords[i, 1], coords[i+1, 1]], \n                [coords[i, 2], coords[i+1, 2]], 'k-', lw=1)\n    ax.set_title(title)\n    ax.set_xlabel('X (Å)')\n    ax.set_ylabel('Y (Å)')\n    ax.set_zlabel('Z (Å)')\n    ax.legend()\n    plt.savefig(f\"{title.replace(' ', '_')}.png\")\n    plt.close()\n\n# 12. (Optional) Ensemble Modeling\nclass ModelEnsemble:\n    def __init__(self, models, weights=None):\n        self.models = models\n        self.weights = weights if weights is not None else [1/len(models)] * len(models)\n    def predict(self, features, seq_lengths=None):\n        all_predictions = []\n        for i, model in enumerate(self.models):\n            model.eval()\n            with torch.no_grad():\n                output = model(features, seq_lengths)\n                all_predictions.append(output * self.weights[i])\n        return sum(all_predictions)\n\n# 13. Main Execution\ndef main():\n    print(\"\\n--- Main execution ---\")\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n    \n    input_dim = train_data[0]['features'].shape[1]\n    model = RNAFoldingModel(input_dim=input_dim).to(device)\n    \n    print(\"\\nStarting model training...\")\n    trained_model = train_model(\n        model=model,\n        train_loader=train_loader,\n        val_loader=validation_loader,\n        epochs=20,  # 使用正确的参数名\n        lr=0.0005,\n        device=device\n    )\n    \n    \n    print(\"\\nGenerating predictions on test data...\")\n    test_predictions = generate_predictions(trained_model, test_loader, device, num_predictions=5)\n    print(\"\\nPredictions generated.\")\n    \n    print(\"\\nCreating submission file...\")\n    submission_file = create_submission_file(test_predictions, test_sequences)\n    print(f\"\\nSubmission file created: submission.csv\")\n    print(submission_file.head())\n    \n    print(\"\\nVisualizing a sample prediction (first test sequence)...\")\n    sample_seq_id = test_sequences['target_id'].iloc[0]\n    if sample_seq_id in test_predictions:\n        sample_prediction = test_predictions[sample_seq_id][0]\n        visualize_3d_structure(sample_prediction, title=f\"Predicted 3D Structure - {sample_seq_id}\")\n        print(f\"Visualization saved for {sample_seq_id}.\")\n    else:\n        print(\"No prediction found for the first test sequence for visualization.\")\n    \n    print(\"\\n--- Main execution completed ---\")\n\nif __name__ == '__main__':\n    main()\n\nprint(\"\\nNotebook execution finished.\")","metadata":{"_uuid":"4a46c934-fe3f-460c-88cf-b1f1a9d250bb","_cell_guid":"2027a68e-272a-4b5d-9711-089e5e0a70da","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-03-18T07:54:03.186603Z","iopub.execute_input":"2025-03-18T07:54:03.187003Z","iopub.status.idle":"2025-03-18T07:55:11.767725Z","shell.execute_reply.started":"2025-03-18T07:54:03.186969Z","shell.execute_reply":"2025-03-18T07:55:11.766620Z"}},"outputs":[],"execution_count":null}]}