{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210},{"sourceType":"datasetVersion","sourceId":14721765,"datasetId":9364067,"databundleVersionId":15568790}],"dockerImageVersionId":31259,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport torch\nimport torch.nn as nn\nimport os\nimport numpy as np\nimport sys\nimport time\nimport math\n\n# ==========================================\n# 1. Model Architecture (同步本地参数: 128/256)\n# ==========================================\nclass PositionalEncoding(nn.Module):\n    def __init__(self, d_model, max_len=5000):\n        super().__init__()\n        pe = torch.zeros(max_len, d_model)\n        position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        self.register_buffer('pe', pe)\n\n    def forward(self, x):\n        batch_size, seq_len, _ = x.size()\n        return x + self.pe[:seq_len, :].unsqueeze(0)\n\nclass SimpleRNAModel(nn.Module):\n    # 【注意】这里默认值改成了 128 和 256，和你本地训练的一致\n    def __init__(self, embed_dim=128, hidden_dim=256, num_layers=2, nhead=4):\n        super().__init__()\n        self.embedding = nn.Embedding(num_embeddings=5, embedding_dim=embed_dim)\n        self.pos_encoder = PositionalEncoding(embed_dim)\n        encoder_layer = nn.TransformerEncoderLayer(d_model=embed_dim, nhead=nhead, dim_feedforward=hidden_dim, batch_first=True)\n        self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)\n        self.head = nn.Linear(embed_dim, 5 * 3)\n        \n    def forward(self, x):\n        x = self.embedding(x)\n        x = self.pos_encoder(x)\n        out = self.transformer_encoder(x)\n        coords = self.head(out)\n        batch_size, length, _ = coords.shape\n        coords = coords.view(batch_size, length, 5, 3)\n        return coords\n\n# ==========================================\n# 2. Config & Physics\n# ==========================================\nTEST_CSV_PATH = '/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv'\nSUBMISSION_PATH = 'submission.csv'\nMODEL_PATH = '/kaggle/input/guetlit-rna-baseline/baseline_model.pth'\nTOKEN_MAP = {'A': 0, 'G': 1, 'C': 2, 'U': 3}\nBOND_DISTANCE_TARGET = 6.0\nBOND_DISTANCE_TOL = 0.5\n\n# 列名逻辑\nCOL_ORDER = ['ID', 'resname', 'resid'] + [f'{axis}_{m}' for m in range(1, 6) for axis in ['x', 'y', 'z']]\n\ndef refine_geometry(coords):\n    refined = coords.copy()\n    length = len(coords)\n    for i in range(length - 1):\n        vec = refined[i+1] - refined[i]\n        dist = np.linalg.norm(vec)\n        if abs(dist - BOND_DISTANCE_TARGET) > BOND_DISTANCE_TOL:\n            target = BOND_DISTANCE_TARGET\n            unit_vec = vec / (dist + 1e-10)\n            refined[i+1] = refined[i] + unit_vec * target\n    return refined\n\n# ==========================================\n# 3. Execution\n# ==========================================\ndef run_kaggle_submission():\n    print(f\"[INFO] Initializing Inference (Dim=128/256)...\")\n    start_time = time.time()\n    \n    if not os.path.exists(MODEL_PATH):\n        print(f\"[ERROR] Model not found at {MODEL_PATH}\")\n        return\n\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    # 实例化模型\n    model = SimpleRNAModel(embed_dim=128, hidden_dim=256)\n    \n    try:\n        state_dict = torch.load(MODEL_PATH, map_location=device)\n        model.load_state_dict(state_dict)\n        model.to(device)\n        model.eval()\n        print(\"[INFO] Weights loaded successfully.\")\n    except Exception as e:\n        print(f\"[ERROR] Weight mismatch or load error: {e}\")\n        return\n\n    chunk_iter = pd.read_csv(TEST_CSV_PATH, chunksize=1000)\n    submission_rows = []\n    \n    print(\"[PROCESS] Prediction loop start...\")\n    \n    with torch.no_grad():\n        for chunk in chunk_iter:\n            for idx, row in chunk.iterrows():\n                seq_id = row['target_id']\n                seq_str = row['sequence']\n                \n                indices = [TOKEN_MAP.get(c, 0) for c in seq_str]\n                input_tensor = torch.LongTensor([indices]).to(device)\n                \n                try:\n                    output = model(input_tensor)\n                    raw_coords = output[0].cpu().numpy()\n                    # 反归一化\n                    raw_coords = raw_coords * 10.0\n                except:\n                    raw_coords = np.zeros((len(seq_str), 5, 3))\n                \n                final_coords = np.zeros_like(raw_coords)\n                for m in range(5):\n                    final_coords[:, m, :] = refine_geometry(raw_coords[:, m, :])\n                \n                length = len(seq_str)\n                for i in range(length):\n                    row_data = {\n                        'ID': f\"{seq_id}_{i+1}\",\n                        'resname': seq_str[i],\n                        'resid': i + 1\n                    }\n                    for m in range(5):\n                        row_data[f'x_{m+1}'] = final_coords[i, m, 0]\n                        row_data[f'y_{m+1}'] = final_coords[i, m, 1]\n                        row_data[f'z_{m+1}'] = final_coords[i, m, 2]\n                    submission_rows.append(row_data)\n\n    df_sub = pd.DataFrame(submission_rows)\n    df_sub = df_sub[COL_ORDER]\n    df_sub.to_csv(SUBMISSION_PATH, index=False)\n    \n    print(f\"[SUCCESS] Submission generated in {time.time()-start_time:.2f}s\")\n\nif __name__ == '__main__':\n    run_kaggle_submission()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-03T17:20:17.70889Z","iopub.execute_input":"2026-02-03T17:20:17.710161Z","iopub.status.idle":"2026-02-03T17:20:28.245944Z","shell.execute_reply.started":"2026-02-03T17:20:17.710092Z","shell.execute_reply":"2026-02-03T17:20:28.244257Z"}},"outputs":[],"execution_count":null}]}