{"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":14928314,"datasetId":9552475,"databundleVersionId":15795629}],"dockerImageVersionId":31286,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"**1. 서열 특성 분석 (Sequence EDA)**\n\nRNA의 길이와 염기(A, C, G, U) 분포를 확인하는 코드","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# 데이터 로드\ntrain_seqs = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/train_sequences.csv')\n\n# 1. 서열 길이 분석\ntrain_seqs['seq_len'] = train_seqs['sequence'].apply(len)\n\nplt.figure(figsize=(10, 5))\nsns.histplot(train_seqs['seq_len'], bins=50, kde=True, color='skyblue')\nplt.title('RNA Sequence Length Distribution')\nplt.xlabel('Length (number of nucleotides)')\nplt.ylabel('Count')\nplt.show()\n\n# 2. 염기(A, C, G, U) 조성비 분석\ndef get_base_content(seq):\n    return {base: seq.count(base) / len(seq) for base in 'ACGU'}\n\nbase_ratios = train_seqs['sequence'].apply(get_base_content).apply(pd.Series)\nprint(\"--- 평균 염기 조성비 ---\")\nprint(base_ratios.mean())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:10:56.841897Z","iopub.execute_input":"2026-02-24T09:10:56.842198Z","iopub.status.idle":"2026-02-24T09:10:57.98484Z","shell.execute_reply.started":"2026-02-24T09:10:56.842177Z","shell.execute_reply":"2026-02-24T09:10:57.98378Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Sequence\n   ↓\nDistance Transformer\n   ↓\nMDS initial coords\n   ↓\nSE(3) refinement\n   ↓\nDiffusion sampling\n   ↓\n5 conformations","metadata":{}},{"cell_type":"markdown","source":"1️⃣ 데이터 로드 셀 (추가)","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\nDATA_PATH = \"/kaggle/input/stanford-rna-3d-folding-2/\"\n\ntrain_seqs = pd.read_csv(DATA_PATH + \"train_sequences.csv\")\ntrain_labels = pd.read_csv(DATA_PATH + \"train_labels.csv\")\n\nprint(\"train_seqs:\", train_seqs.shape)\nprint(\"train_labels:\", train_labels.shape)\ntrain_seqs.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:10:57.986274Z","iopub.execute_input":"2026-02-24T09:10:57.986585Z","iopub.status.idle":"2026-02-24T09:11:03.540829Z","shell.execute_reply.started":"2026-02-24T09:10:57.986554Z","shell.execute_reply":"2026-02-24T09:11:03.539369Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Dataset 생성","metadata":{}},{"cell_type":"markdown","source":"📦 1️⃣ Distogram Dataset","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport pandas as pd\nimport numpy as np\n\n# -------------------------\n# Config\n# -------------------------\nMIN_DIST = 2.0\nMAX_DIST = 50.0\nNUM_BINS = 64\n\nBIN_WIDTH = (MAX_DIST - MIN_DIST) / (NUM_BINS - 1)\n\nBASE2IDX = {\"A\":0, \"C\":1, \"G\":2, \"U\":3}\n\n# -------------------------\n# Helper: distance → bin\n# -------------------------\ndef dist_to_bins(dist_matrix):\n    bins = ((dist_matrix - MIN_DIST) / BIN_WIDTH).long()\n    bins = torch.clamp(bins, 0, NUM_BINS - 1)\n    return bins\n\n# -------------------------\n# Dataset\n# -------------------------\nclass RNADistogramDataset(torch.utils.data.Dataset):\n    def __init__(self, seq_df, label_df):\n        self.seq_df = seq_df\n        self.label_df = label_df\n        \n        # build coord dict\n        self.coords = {}\n        prefixes = label_df[\"ID\"].str.rsplit(\"_\", n=1).str[0]\n        for pid, group in label_df.groupby(prefixes):\n            self.coords[pid] = group.sort_values(\"resid\")[[\"x_1\",\"y_1\",\"z_1\"]].values\n\n    def __len__(self):\n        return len(self.seq_df)\n\n    def encode_seq(self, seq):\n        idx = torch.tensor([BASE2IDX[b] for b in seq])\n        onehot = F.one_hot(idx, num_classes=4).float()\n        return onehot\n\n    def __getitem__(self, idx):\n        row = self.seq_df.iloc[idx]\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n\n        node_feat = self.encode_seq(seq)   # (L,4)\n        coords = torch.tensor(self.coords[tid], dtype=torch.float)\n\n        dist = torch.cdist(coords, coords)   # (L,L)\n        bins = dist_to_bins(dist)\n\n        return {\n            \"node_feat\": node_feat,\n            \"dist_bins\": bins,\n            \"mask\": torch.ones(len(seq), dtype=torch.bool)\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:03.541794Z","iopub.execute_input":"2026-02-24T09:11:03.542079Z","iopub.status.idle":"2026-02-24T09:11:03.552698Z","shell.execute_reply.started":"2026-02-24T09:11:03.542051Z","shell.execute_reply":"2026-02-24T09:11:03.551926Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"🧠 2️⃣ Model — 1D + 2D Transformer","metadata":{}},{"cell_type":"code","source":"# -------------------------\n# 1D Encoder\n# -------------------------\nclass SequenceEncoder(nn.Module):\n    def __init__(self, d_model=256, nhead=8, layers=3):\n        super().__init__()\n        self.input_proj = nn.Linear(4, d_model)\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model,\n            nhead=nhead,    # 여기서 nhead를 d_model과 맞추기\n            batch_first=True\n        )\n        self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=layers)\n\n    def forward(self, x):\n        x = self.input_proj(x)\n        return self.encoder(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:03.554326Z","iopub.execute_input":"2026-02-24T09:11:03.554659Z","iopub.status.idle":"2026-02-24T09:11:03.577711Z","shell.execute_reply.started":"2026-02-24T09:11:03.554634Z","shell.execute_reply":"2026-02-24T09:11:03.576519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------\n# Pair Initialization\n# -------------------------\nclass PairInit(nn.Module):\n    def __init__(self, d_node=256, d_pair=128):\n        super().__init__()\n        self.proj = nn.Linear(d_node*2, d_pair)\n\n    def forward(self, node_repr):\n        L = node_repr.size(1)\n        a = node_repr.unsqueeze(2).expand(-1, L, L, -1)\n        b = node_repr.unsqueeze(1).expand(-1, L, L, -1)\n        pair = torch.cat([a,b], dim=-1)\n        return self.proj(pair)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:03.578737Z","iopub.execute_input":"2026-02-24T09:11:03.578984Z","iopub.status.idle":"2026-02-24T09:11:03.598992Z","shell.execute_reply.started":"2026-02-24T09:11:03.578965Z","shell.execute_reply":"2026-02-24T09:11:03.59824Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------\n# 2D Pair Block\n# -------------------------\nclass PairBlock(nn.Module):\n    def __init__(self, d_pair=128, nhead=8):\n        super().__init__()\n        self.row_attn = nn.MultiheadAttention(d_pair, nhead, batch_first=True)\n        self.col_attn = nn.MultiheadAttention(d_pair, nhead, batch_first=True)\n        self.ff = nn.Sequential(\n            nn.Linear(d_pair, d_pair*4),\n            nn.ReLU(),\n            nn.Linear(d_pair*4, d_pair)\n        )\n        self.norm1 = nn.LayerNorm(d_pair)\n        self.norm2 = nn.LayerNorm(d_pair)\n        self.norm3 = nn.LayerNorm(d_pair)\n\n    def forward(self, x):\n        B, L, _, C = x.shape\n        \n        # row attention\n        x_row = x.reshape(B*L, L, C)\n        attn_out,_ = self.row_attn(x_row, x_row, x_row)\n        x_row = self.norm1(x_row + attn_out)\n        x = x_row.reshape(B, L, L, C)\n        \n        # column attention\n        x_col = x.transpose(1,2).reshape(B*L, L, C)\n        attn_out,_ = self.col_attn(x_col, x_col, x_col)\n        x_col = self.norm2(x_col + attn_out)\n        x = x_col.reshape(B, L, L, C).transpose(1,2)\n\n        # feedforward\n        x = self.norm3(x + self.ff(x))\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:03.599873Z","iopub.execute_input":"2026-02-24T09:11:03.600064Z","iopub.status.idle":"2026-02-24T09:11:03.612844Z","shell.execute_reply.started":"2026-02-24T09:11:03.600045Z","shell.execute_reply":"2026-02-24T09:11:03.611603Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------\n# Full Model\n# -------------------------\nclass RNADistogramModel(nn.Module):\n    def __init__(self, d_node=256, d_pair=128, pair_layers=6, nhead=8):  # nhead 추가\n        super().__init__()\n        self.seq_encoder = SequenceEncoder(d_node, nhead=nhead)  # nhead 전달\n        self.pair_init = PairInit(d_node, d_pair)\n        self.pair_blocks = nn.ModuleList([\n            PairBlock(d_pair) for _ in range(pair_layers)\n        ])\n        self.head = nn.Linear(d_pair, NUM_BINS)\n\n    def forward(self, node_feat):\n        node = self.seq_encoder(node_feat)\n        pair = self.pair_init(node)\n        for block in self.pair_blocks:\n            pair = block(pair)\n\n        logits = self.head(pair)\n        logits = (logits + logits.transpose(1,2)) / 2  # symmetry\n        return logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:15:06.319011Z","iopub.execute_input":"2026-02-24T09:15:06.319268Z","iopub.status.idle":"2026-02-24T09:15:06.325841Z","shell.execute_reply.started":"2026-02-24T09:15:06.319248Z","shell.execute_reply":"2026-02-24T09:15:06.324911Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"🏋️ 3️⃣ Training Loop","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, device):\n    model.train()\n    total_loss = 0\n\n    for batch in loader:\n        node = batch[\"node_feat\"].to(device)\n        target = batch[\"dist_bins\"].to(device)\n\n        optimizer.zero_grad()\n        logits = model(node)\n\n        B,L,L2,BINS = logits.shape\n        loss = F.cross_entropy(\n            logits.view(-1, NUM_BINS),\n            target.view(-1)\n        )\n\n        loss.backward()\n        optimizer.step()\n        total_loss += loss.item()\n\n    return total_loss / len(loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:03.628079Z","iopub.execute_input":"2026-02-24T09:11:03.628368Z","iopub.status.idle":"2026-02-24T09:11:03.63694Z","shell.execute_reply.started":"2026-02-24T09:11:03.628347Z","shell.execute_reply":"2026-02-24T09:11:03.636006Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"✅ 1️⃣ Collate Function 구현","metadata":{}},{"cell_type":"code","source":"def rna_collate_fn(batch):\n    \"\"\"\n    batch: list of dict\n        node_feat: (L,4)\n        dist_bins: (L,L)\n        mask: (L,)\n    \"\"\"\n\n    max_len = max(item[\"node_feat\"].size(0) for item in batch)\n\n    node_feats = []\n    dist_bins = []\n    masks = []\n\n    for item in batch:\n        L = item[\"node_feat\"].size(0)\n        pad_len = max_len - L\n\n        # Node padding\n        node_pad = F.pad(item[\"node_feat\"], (0,0,0,pad_len))  # (max_len,4)\n        node_feats.append(node_pad)\n\n        # Distance padding\n        dist_pad = F.pad(item[\"dist_bins\"],\n                         (0,pad_len,0,pad_len),\n                         value=0)\n        dist_bins.append(dist_pad)\n\n        # Mask padding\n        mask_pad = F.pad(item[\"mask\"],\n                         (0,pad_len),\n                         value=0)\n        masks.append(mask_pad)\n\n    return {\n        \"node_feat\": torch.stack(node_feats),      # (B, Lmax, 4)\n        \"dist_bins\": torch.stack(dist_bins),      # (B, Lmax, Lmax)\n        \"mask\": torch.stack(masks)                # (B, Lmax)\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:03.637912Z","iopub.execute_input":"2026-02-24T09:11:03.638213Z","iopub.status.idle":"2026-02-24T09:11:03.655953Z","shell.execute_reply.started":"2026-02-24T09:11:03.638112Z","shell.execute_reply":"2026-02-24T09:11:03.655055Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset = RNADistogramDataset(train_seqs, train_labels)\n\nloader = torch.utils.data.DataLoader(\n    dataset,\n    batch_size=1,\n    shuffle=True,\n    collate_fn=rna_collate_fn\n)\n\nprint(\"Dataset size:\", len(dataset))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:03.65829Z","iopub.execute_input":"2026-02-24T09:11:03.658555Z","iopub.status.idle":"2026-02-24T09:11:13.642761Z","shell.execute_reply.started":"2026-02-24T09:11:03.658535Z","shell.execute_reply":"2026-02-24T09:11:13.641556Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch = next(iter(loader))\n\nprint(\"node_feat:\", batch[\"node_feat\"].shape)\nprint(\"dist_bins:\", batch[\"dist_bins\"].shape)\nprint(\"mask:\", batch[\"mask\"].shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:13.643907Z","iopub.execute_input":"2026-02-24T09:11:13.644193Z","iopub.status.idle":"2026-02-24T09:11:13.681757Z","shell.execute_reply.started":"2026-02-24T09:11:13.644172Z","shell.execute_reply":"2026-02-24T09:11:13.680922Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"✅ 2️⃣ DataLoader 적용","metadata":{}},{"cell_type":"markdown","source":"✅ 3️⃣ Model Forward 수정 (mask 적용 준비)","metadata":{}},{"cell_type":"markdown","source":"✅ 4️⃣ Padding-aware Loss","metadata":{}},{"cell_type":"code","source":"def masked_distogram_loss(logits, target, mask):\n    \"\"\"\n    logits: (B,L,L,BINS)\n    target: (B,L,L)\n    mask:   (B,L)\n    \"\"\"\n\n    B,L,_,_ = logits.shape\n\n    # pair mask 생성\n    pair_mask = mask.unsqueeze(1) & mask.unsqueeze(2)  # (B,L,L)\n\n    # diagonal 제거\n    eye = torch.eye(L, device=mask.device).bool()\n    pair_mask = pair_mask & (~eye.unsqueeze(0))\n\n    logits = logits[pair_mask]\n    target = target[pair_mask]\n\n    return F.cross_entropy(logits, target)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:13.682594Z","iopub.execute_input":"2026-02-24T09:11:13.682784Z","iopub.status.idle":"2026-02-24T09:11:13.688308Z","shell.execute_reply.started":"2026-02-24T09:11:13.682766Z","shell.execute_reply":"2026-02-24T09:11:13.687319Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"✅ 5️⃣ Training Loop 수정","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, device):\n    model.train()\n    total_loss = 0\n\n    for batch in loader:\n        node = batch[\"node_feat\"].to(device)\n        target = batch[\"dist_bins\"].to(device)\n        mask = batch[\"mask\"].to(device)\n\n        optimizer.zero_grad()\n        logits = model(node)\n\n        loss = masked_distogram_loss(logits, target, mask)\n\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n\n    return total_loss / len(loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:13.689644Z","iopub.execute_input":"2026-02-24T09:11:13.689959Z","iopub.status.idle":"2026-02-24T09:11:13.706427Z","shell.execute_reply.started":"2026-02-24T09:11:13.689928Z","shell.execute_reply":"2026-02-24T09:11:13.705565Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"1️⃣ Distogram → 기대 거리 계산","metadata":{}},{"cell_type":"code","source":"import torch\n\ndef distogram_to_expected_distance(dist_logits):\n    \"\"\"\n    dist_logits: (L, L, B) or (B, L, L, B)\n    return: expected_dist: (L,L)\n    \"\"\"\n    if dist_logits.dim() == 4:\n        # batch 처리\n        probs = torch.softmax(dist_logits, dim=-1)\n        bin_centers = torch.linspace(\n            MIN_DIST + BIN_WIDTH/2, \n            MAX_DIST - BIN_WIDTH/2, \n            NUM_BINS, device=dist_logits.device\n        )\n        expected_dist = (probs * bin_centers).sum(-1)\n        return expected_dist\n    else:\n        probs = torch.softmax(dist_logits, dim=-1)\n        bin_centers = torch.linspace(\n            MIN_DIST + BIN_WIDTH/2, \n            MAX_DIST - BIN_WIDTH/2, \n            NUM_BINS, device=dist_logits.device\n        )\n        expected_dist = (probs * bin_centers).sum(-1)\n        return expected_dist","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:13.707724Z","iopub.execute_input":"2026-02-24T09:11:13.708124Z","iopub.status.idle":"2026-02-24T09:11:13.724831Z","shell.execute_reply.started":"2026-02-24T09:11:13.708099Z","shell.execute_reply":"2026-02-24T09:11:13.723656Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"2️⃣ Classical MDS로 초기 좌표 복원","metadata":{}},{"cell_type":"code","source":"def mds_from_distance(dist_matrix, n_dim=3):\n    \"\"\"\n    Classical MDS\n    dist_matrix: (L,L)\n    return: coords: (L,3)\n    \"\"\"\n    L = dist_matrix.size(0)\n    H = torch.eye(L, device=dist_matrix.device) - 1.0/L\n    D2 = dist_matrix**2\n    B = -0.5 * H @ D2 @ H\n\n    # Eigen decomposition\n    eigvals, eigvecs = torch.linalg.eigh(B)\n    idx = torch.argsort(eigvals, descending=True)\n    eigvals = eigvals[idx][:n_dim]\n    eigvecs = eigvecs[:, idx][:, :n_dim]\n\n    coords = eigvecs * torch.sqrt(eigvals)\n    return coords","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:13.725755Z","iopub.execute_input":"2026-02-24T09:11:13.726054Z","iopub.status.idle":"2026-02-24T09:11:13.736573Z","shell.execute_reply.started":"2026-02-24T09:11:13.726035Z","shell.execute_reply":"2026-02-24T09:11:13.735707Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"예시: Distogram → 초기 좌표","metadata":{}},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nmodel = RNADistogramModel(\n    d_node=256,\n    d_pair=128,\n    pair_layers=4   # 메모리 상황에 맞게 조절\n).to(device)\n\n# Optimizer 예시\nimport torch.optim as optim\noptimizer = optim.Adam(model.parameters(), lr=1e-3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:13.737769Z","iopub.execute_input":"2026-02-24T09:11:13.737982Z","iopub.status.idle":"2026-02-24T09:11:13.780869Z","shell.execute_reply.started":"2026-02-24T09:11:13.737964Z","shell.execute_reply":"2026-02-24T09:11:13.779771Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample = next(iter(loader))\nnode_feat = sample[\"node_feat\"][0].to(device)       # (Lmax, 4)\nmask = sample[\"mask\"][0].to(device)\n\nmodel.eval()\nwith torch.no_grad():\n    logits = model(node_feat.unsqueeze(0))    # (1, Lmax, Lmax, B)\n    L = mask.sum()\n    logits = logits[0, :L, :L]                # 패딩 제거\n    expected_dist = distogram_to_expected_distance(logits)\n    coords_init = mds_from_distance(expected_dist)\n\nprint(\"Initial coords shape:\", coords_init.shape)   # (L,3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:13.782Z","iopub.execute_input":"2026-02-24T09:11:13.782356Z","iopub.status.idle":"2026-02-24T09:11:13.990833Z","shell.execute_reply.started":"2026-02-24T09:11:13.782324Z","shell.execute_reply":"2026-02-24T09:11:13.990221Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"1️⃣ EGNN 스타일 SE(3)-Equivariant Block","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\nclass EGNNLayer(nn.Module):\n    def __init__(self, node_dim, edge_dim=0):\n        super().__init__()\n        self.edge_mlp = nn.Sequential(\n            nn.Linear(2*node_dim + 1 + edge_dim, 128),\n            nn.ReLU(),\n            nn.Linear(128, 1)\n        )\n        self.node_mlp = nn.Sequential(\n            nn.Linear(node_dim + 1, 128),  # node_dim + aggregated distance\n            nn.ReLU(),\n            nn.Linear(128, node_dim)\n        )\n\n    def forward(self, h, x, mask=None):\n        B, L, _ = x.shape\n        dx = x.unsqueeze(2) - x.unsqueeze(1)\n        dist = torch.norm(dx, dim=-1, keepdim=True)\n\n        h_i = h.unsqueeze(2).expand(-1, -1, L, -1)\n        h_j = h.unsqueeze(1).expand(-1, L, -1, -1)\n\n        edge_input = torch.cat([h_i, h_j, dist], dim=-1)\n        e_ij = self.edge_mlp(edge_input)\n\n        delta_x = (dx / (dist + 1e-8)) * e_ij\n        x = x + delta_x.sum(dim=2)\n\n        # node update\n        agg_dist = dist.mean(dim=2)       # (B,L,1)\n        agg = torch.cat([h, agg_dist], dim=-1)\n        h = h + self.node_mlp(agg)\n\n        if mask is not None:\n            h = h * mask.unsqueeze(-1)\n            x = x * mask.unsqueeze(-1)\n\n        return h, x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:13.991705Z","iopub.execute_input":"2026-02-24T09:11:13.991947Z","iopub.status.idle":"2026-02-24T09:11:13.999019Z","shell.execute_reply.started":"2026-02-24T09:11:13.991919Z","shell.execute_reply":"2026-02-24T09:11:13.998174Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"2️⃣ SE(3)-Equivariant Refinement Network","metadata":{}},{"cell_type":"code","source":"class EGNNRefine(nn.Module):\n    def __init__(self, node_dim=256, n_layers=4):\n        super().__init__()\n        self.input_proj = nn.Linear(4, node_dim)  # RNA one-hot → embedding\n        self.layers = nn.ModuleList([EGNNLayer(node_dim) for _ in range(n_layers)])\n    \n    def forward(self, node_feat, coords_init, mask=None):\n        \"\"\"\n        node_feat: (B, L, 4)\n        coords_init: (B, L, 3)\n        mask: (B, L)\n        \"\"\"\n        h = self.input_proj(node_feat)\n        x = coords_init\n        for layer in self.layers:\n            h, x = layer(h, x, mask)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:14.000255Z","iopub.execute_input":"2026-02-24T09:11:14.0006Z","iopub.status.idle":"2026-02-24T09:11:14.019223Z","shell.execute_reply.started":"2026-02-24T09:11:14.000565Z","shell.execute_reply":"2026-02-24T09:11:14.017913Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"3️⃣ 사용 예시","metadata":{}},{"cell_type":"code","source":"# 단일 배치 테스트\nsample = next(iter(loader))\nnode_feat = sample[\"node_feat\"].to(device)   # (B,L,4)\nmask = sample[\"mask\"].to(device)\n\n# MDS 초기 좌표 계산\nmodel.eval()\nwith torch.no_grad():\n    logits = model(node_feat)\n    expected_dist = distogram_to_expected_distance(logits)\n    coords_init = torch.stack([mds_from_distance(ed) for ed in expected_dist])  # (B,L,3)\n\n# SE(3) refinement\negnn_model = EGNNRefine(node_dim=256, n_layers=4).to(device)\ncoords_refined = egnn_model(node_feat, coords_init, mask)\n\nprint(\"Refined coords shape:\", coords_refined.shape)  # (B,L,3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:14.020744Z","iopub.execute_input":"2026-02-24T09:11:14.020991Z","iopub.status.idle":"2026-02-24T09:11:14.11037Z","shell.execute_reply.started":"2026-02-24T09:11:14.020966Z","shell.execute_reply":"2026-02-24T09:11:14.109649Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"1️⃣ Multi-structure diffusion head (간단 버전)","metadata":{}},{"cell_type":"code","source":"class MultiStructureHead(nn.Module):\n    \"\"\"\n    Refined coordinates를 기반으로 5개의 후보 구조를 생성\n    - noise injection + learned offsets\n    \"\"\"\n    def __init__(self, n_structures=5, node_dim=256):\n        super().__init__()\n        self.n_structures = n_structures\n        self.offset_mlp = nn.Sequential(\n            nn.Linear(node_dim + 3, 128),\n            nn.ReLU(),\n            nn.Linear(128, 3)\n        )\n\n    def forward(self, h, coords_refined):\n        \"\"\"\n        h: node embedding (B,L,node_dim)\n        coords_refined: (B,L,3)\n        \"\"\"\n        B, L, _ = coords_refined.shape\n        outputs = []\n\n        for i in range(self.n_structures):\n            # noise injection\n            noise = torch.randn_like(coords_refined) * 0.1 * (i+1)\n            inp = torch.cat([h, coords_refined + noise], dim=-1)\n            offset = self.offset_mlp(inp)\n            coords_out = coords_refined + offset\n            outputs.append(coords_out)\n\n        return torch.stack(outputs, dim=1)   # (B, n_structures, L, 3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:14.111399Z","iopub.execute_input":"2026-02-24T09:11:14.111672Z","iopub.status.idle":"2026-02-24T09:11:14.118028Z","shell.execute_reply.started":"2026-02-24T09:11:14.11165Z","shell.execute_reply":"2026-02-24T09:11:14.117072Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"2️⃣ Pipeline 통합","metadata":{}},{"cell_type":"code","source":"# 1) Distogram → MDS\nsample = next(iter(loader))\nnode_feat = sample[\"node_feat\"].to(device)\nmask = sample[\"mask\"].to(device)\n\nmodel.eval()\nwith torch.no_grad():\n    logits = model(node_feat)                         # (B,L,L,B)\n    expected_dist = distogram_to_expected_distance(logits)\n    coords_init = torch.stack([mds_from_distance(ed) for ed in expected_dist])\n\n# 2) SE(3)-Equivariant refinement\negnn_model = EGNNRefine(node_dim=256, n_layers=4).to(device)\nh = egnn_model.input_proj(node_feat)                 # node embedding\ncoords_refined = egnn_model(node_feat, coords_init, mask)  # (B,L,3)\n\n# 3) Multi-structure head → 5 candidate structures\nmulti_head = MultiStructureHead(n_structures=5, node_dim=256).to(device)\ncoords_5 = multi_head(h, coords_refined)            # (B,5,L,3)\n\nprint(\"Final 5 structures shape:\", coords_5.shape)  # (B,5,L,3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:11:14.118947Z","iopub.execute_input":"2026-02-24T09:11:14.119174Z","iopub.status.idle":"2026-02-24T09:11:14.186792Z","shell.execute_reply.started":"2026-02-24T09:11:14.119156Z","shell.execute_reply":"2026-02-24T09:11:14.185609Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"1️⃣ Submission 변환 코드","metadata":{}},{"cell_type":"code","source":"d_node = 64\nnhead = 8\nmodel = RNADistogramModel(d_node=d_node, d_pair=128, pair_layers=3, nhead=nhead).to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:15:57.510524Z","iopub.execute_input":"2026-02-24T09:15:57.510814Z","iopub.status.idle":"2026-02-24T09:15:57.529703Z","shell.execute_reply.started":"2026-02-24T09:15:57.510795Z","shell.execute_reply":"2026-02-24T09:15:57.528707Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# RNA 3D Folding: Distogram → MDS → submission.csv\n# ============================================================\n\nimport pandas as pd\nimport torch\nimport numpy as np\nfrom sklearn.manifold import MDS\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nDATA_PATH = \"/kaggle/input/stanford-rna-3d-folding-2/\"\n\n# -----------------------------\n# 1) 데이터 로드\n# -----------------------------\ntrain_seqs = pd.read_csv(DATA_PATH + \"train_sequences.csv\")\ntest_seqs  = pd.read_csv(DATA_PATH + \"test_sequences.csv\")\nprint(\"Train:\", train_seqs.shape, \"Test:\", test_seqs.shape)\n\n# -----------------------------\n# 2) 모델 정의 (예시)\n# -----------------------------\n# 이미 정의되어 있다고 가정\n# RNADistogramModel(d_node, d_pair, pair_layers)\n\n# 노드 feature dimension: one-hot A,C,G,U → 4\ndef one_hot_encode(seq):\n    mapping = {'A':0,'C':1,'G':2,'U':3}\n    L = len(seq)\n    feat = np.zeros((L, 4), dtype=np.float32)\n    for i, s in enumerate(seq):\n        feat[i, mapping[s]] = 1.0\n    return feat\n\n# -----------------------------\n# 3) distogram → MDS → 5개 좌표 생성\n# -----------------------------\nall_predictions = []\n\n# 모델 초기화\nmodel = RNADistogramModel(d_node=4, d_pair=128, pair_layers=3).to(device)\nmodel.eval()\n\nfor idx, row in test_seqs.iterrows():\n    tid = row[\"target_id\"]\n    seq = row[\"sequence\"]\n    L = len(seq)\n    \n    # node feature\n    node_feat = torch.tensor(one_hot_encode(seq), device=device).unsqueeze(0)  # (1,L,4)\n    \n    with torch.no_grad():\n        dist_pred = model(node_feat)  # (1,L,L) or (1,L,L,B)\n        dist_pred = dist_pred[0].cpu().numpy()  # (L,L)\n    \n    # MDS로 좌표 재구성\n    mds = MDS(n_components=3, dissimilarity='precomputed', random_state=42)\n    coords_base = mds.fit_transform(dist_pred)  # (L,3)\n    \n    # 5개 구조 생성 (조금씩 랜덤 perturb)\n    coords_5 = []\n    rng = np.random.default_rng(42)\n    for i in range(5):\n        coords_5.append(coords_base + rng.normal(scale=0.5, size=coords_base.shape))\n    coords_5 = np.stack(coords_5, axis=0)  # (5,L,3)\n    \n    # submission row 생성\n    for j in range(L):\n        res = {\"ID\": f\"{tid}_{j+1}\", \"resname\": seq[j], \"resid\": j+1}\n        for k in range(5):\n            res[f\"x_{k+1}\"], res[f\"y_{k+1}\"], res[f\"z_{k+1}\"] = coords_5[k,j]\n        all_predictions.append(res)\n\n# -----------------------------\n# 4) CSV 생성\n# -----------------------------\nsub = pd.DataFrame(all_predictions)\ncols = [\"ID\",\"resname\",\"resid\"] + [f\"{c}_{i}\" for i in range(1,6) for c in [\"x\",\"y\",\"z\"]]\nsub = sub[cols]\n\n# Kaggle 기준 clipping\ncoord_cols = [c for c in cols if c.startswith((\"x_\",\"y_\",\"z_\"))]\nsub[coord_cols] = sub[coord_cols].clip(-999.999, 9999.999)\n\nsub.to_csv(\"submission.csv\", index=False)\nprint(\"submission.csv 생성 완료!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T09:15:59.042122Z","iopub.execute_input":"2026-02-24T09:15:59.042716Z","iopub.status.idle":"2026-02-24T09:15:59.37114Z","shell.execute_reply.started":"2026-02-24T09:15:59.042659Z","shell.execute_reply":"2026-02-24T09:15:59.369441Z"}},"outputs":[],"execution_count":null}]}